Author SHA1 Message Date
blepping d8aadb2359 Perlin documentation and cleanups
Add OCSNoise PerlinSimple node
2024-08-27 07:52:12 -06:00
blepping 9b7e6b6083 Enable alt_cfgpp_scale for dpmpp_2s, dpmpp_sde and res samplers 2024-08-27 04:59:06 -06:00
blepping 7fc26d0488 Euler should allow alt_cfgpp 2024-08-25 20:27:00 -06:00
blepping 7418974f6b Initial Perlin3D and 2D implementation 2024-08-21 06:50:54 -06:00
43 changed files with 4120 additions and 11001 deletions
+66 -308
View File
@@ -6,38 +6,31 @@ Experimental and mathematically unsound (but fun!) sampling for [ComfyUI](https:
Feel free create a question in Discussions for usage help: [OCS Q&A Discussion](https://github.com/blepping/comfyui_overly_complicated_sampling/discussions/categories/q-a)
_Note for Flux users_: Set `cfg1_uncond_optimization: true` in the `model` block for the main `OCS Sampler` as Flux does not use CFG. CFG++ and alt CFG++ features do not work with Flux.
## Features
* Many different samplers.
* Allows scheduling samplers (i.e. run `euler` for steps 1-4, then switch to `dpmpp_sde`).
* CFG++ support (for some samplers).
* Native support for Restart sigmas. (Restarts do not currently work with RF models like Flux.)
* Native support for Restart sigmas.
* Supports custom noise types.
* Immiscible noise for sampling and Restart. See https://arxiv.org/abs/2406.12303 (note that it was designed for training not inference).
* Allows splitting/combining steps in various ways for (potentially) more accurate sampling.
* Supports Diffrax, torchdiffeq, torchode and torchsde solver backends. (SDE mode not recommended currently.)
* Many tuneable parameters to play with.
* Supports ancestral sampling (in a janky way) for rectified flow models like Flux, works for most basic samplers: does not work for SDE samplers currently.
* Built in safe expression language that allows filtering and manipulating nearly all parameters during sampling.
## Credits
I can move code around but sampling math and creating samplers generally beyond my ability. I didn't write any of the original samplers:
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:
* 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).
* 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 and Euler Dancing samplers based on implementation from 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 https://github.com/yhli123/Immiscible-Diffusion - idea for sampling with it and implementation help from https://github.com/Clybius
* 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
* 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
@@ -98,9 +91,6 @@ 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.
@@ -110,17 +100,13 @@ restart:
# The noise block allows defining global noise sampling parameters.
noise:
# You can disable this to allow GPU noise generation.
# You can disable this to allow GPU noise generation. I believe it only makes a difference for Brownian.
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: 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
set_seed: false
# Global scale scale for generated noise
scale: 1.0
@@ -137,10 +123,10 @@ noise:
# When caching, the batch size for chunks of noise to generate in advance. Generating a batch of noise
# can be more efficient than generating on demand when using a high number of substeps (>10) per step.
batch_size: 1
batch_size: 32
# Whether to cache noise.
caching: true
caching: false
# Interval (in full steps) to reset the cache. Brownian noise takes time into account so
# if using Brownian with caching enabled you will generally want to reset each step.
@@ -187,31 +173,30 @@ 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
```
@@ -239,25 +224,15 @@ When running multiple substeps per step, the results will combined based on the
* `simple`: Doesn't merge anything: only runs a single substep per step.
* `divide`: Creates a linear schedule between the current sigma and the next and runs the substeps in sequence. The model is called at least once per substep.
* `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`).
* `normal`: 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.
<!--
* `average`: The model is called once at the beginning of the step and substeps share that result (but it may be called additional times for higher order samplers). This means substeps for samplers like reversible Euler, Heun 1s, DPM++ 2m SDE are essentially free. May be theoretically very unsound and inaccurate, requires manual tweaking of settings like `s_noise`. Supports the parameter `avgmerge_stretch`(`0.4`) which basically rolls back the current sigma and adds some noise (otherwise running a substep is deterministic and there would be no point to running a sampler like Euler more than once).
* `sample`: Like `average` (and uses `avgmerge_stretch`) but instead of simply using the average, it does a sampler step toward that instead. You can plug in any substep sampler to the `merge_sampler_opt` input (if unconnected and the merge method is `sample` then Euler will be used). *Note*: Substeps in the attached sampler will be ignored.
* `sample_uncached`: Similar to `sample`, however it calls the model per substep instead of caching the result and sharing it. Aside from sampling toward the result, it works more like the `normal` merge strategy. Theoretically it should be better because it's taking less shortcuts but results seem worse.
**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:
```yaml
# Use the simple merge method when step < 3, otherwise use overshoot
dynamic:
- when: step < 3
expression: dict(merge_method :> 'simple)
- expression: dict(merge_method :> 'overshoot)
# You could also write it like this:
dynamic: |
dict(merge_method :> (step < 3 ? 'simple : 'overshoot)
```
When using `average` and `sample` merge strategies and with model call caching enabled you can get away with setting substeps super high. Running something like 100 substeps is actually quite practical and seems to work well.
-->
#### Node Parameters
@@ -280,7 +255,6 @@ The left side group matches steps 0, 1, 2. The right side group matches all step
-->
* `restart_custom_noise`: Currently only used by the `overshoot` merge method.
* `custom_noise`: Currently used by the `lookahead` merge method.
#### Text Parameters
@@ -304,8 +278,6 @@ reta: 1.0
# cond: Shows the model cond prediction (basically the positive prompt).
# uncond: Shows the model uncond prediction (basically the negative prompt).
# raw: Shows the raw noisy latent input.
# noisy: 10% of the noise + denoised.
# diff: Multiplies the difference between cond and uncond.
preview_mode: denoised
# Expression.
@@ -323,26 +295,6 @@ restart:
immiscible:
size: 0
# Only used by the lookahead merge method currently.
lookahead:
# Works like normal samplers, essentially. Disabled by default.
eta: 0.0
# Scales the noise added by lookahead sampling.
s_noise: 1.0
# Controls how much noise to remove in the prediction phase. Higher values will remove more noise.
dt_factor: 1.0
# Immiscible block same as described above.
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
@@ -358,44 +310,33 @@ 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.
* `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`.
* `adapter`: Wraps a normal ComfyUI `SAMPLER`. (Attach a `SAMPLER` parameter to the node.)
* `bogacki`:
* `deis`: See parameters: `history_limit`
* `dpmpp_2m_sde`: 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_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`: If samplers came in vanilla.
* `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`
* `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`:
* `heun`: Alternate Heun implementation. Supports reversible parameters. See parameters: `history_limit`
* `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_1s`: Reversible variant of Heun 1 step. See parameters: `history_limit`.
* `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.
* `heunpp`: See parameters: `max_order`
* `ipndm_v`: See parameters: `history_limit`
* `ipndm`: See parameters: `history_limit`
* `res`
* `reversible_bogacki`:
* `reversible_heun`:
* `reversible_heun_1s`: See parameters: `history_limit`
* `rk4`:
* `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_cycle`: See parameters: `cycle_pct`.
* `trapezoidal`:
* `trapezoidal_cycle`: See parameters: `cycle_pct`
* `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**
@@ -403,46 +344,36 @@ 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)||||
|`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|
|`extraltodeus_distance`|variable|||||
|`gradient_estimation`|1|1||||
|`heun_1s`|1|1||X||
|`heun`|2|||X||
|`heun_1s`|1|1||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_1s`|1|1||X||
|`reversible_heun`|2|||X||
|`rk4`|1-4|||||
|`reversible_heun_1s`|1|1||X||
|`rk4`|4|||||
|`rkf45`|5|||||
|`solver_diffrax`|variable|||||
|`solver_torchdiffeq`|variable|||||
|`solver_torchode`|variable|||||
|`solver_torchsde`|variable|||||
|`trapezoidal`|2|||||
|`trapezoidal_cycle`|2|||||
|`ttm_jvp`|2|||||
`deis`, `ipndm*` and `gradient_estimation` do not seem to work well with ancestralness, I recommend `eta: 0.25` or disable it completely.
`deis`, `ipndm*` do not seem to work well with ancestralness, I recommend `eta: 0.25` or disable it completely.
**Solver Backend Samplers**:
@@ -465,23 +396,6 @@ When doing ancestral sampling, we actually _overshoot_ expected noise for the ne
The difference with cycle is that instead of adding `noise * expected_noise_at_next_step`, we instead first add `noise * (expected_noise_at_next_step * (1.0 - cycle_pct))` and then we generate noise and scale it to `cycle_pct` and add it too. Just for example, suppose `cycle_pct` is `0.2`: we'll add 80% of the expected noise at the next step (`1.0 - 0.2 == 0.8`) and then generate the remaining 20% and add it in to meet the expected amount. I don't recommend setting `cycle_pct` to values over `0.5`, especially if using "weird" noise types.
**Dynamic Step Method**: When `step_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 first matching item will be used. 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). Example:
```yaml
# Use the rk4 merge method when step < 3, otherwise use euler
dynamic:
- when: step < 3
expression: dict(step_method :> 'rk4)
- expression: dict(step_method :> 'euler)
# You could also write it like this:
dynamic: |
dict(step_method :> (step < 3 ? 'rk4 : 'euler)
```
*Note*: You can set all sampler parameters except `substeps` this way. Sampling will use the `substeps` value from the `OCS Substeps` node.
#### Node Parameters
* `substeps`(`1`): Number of substeps. Generally involves a model call per substep, so for example setting this to 4 would approximately quadruple sampling time.
@@ -504,13 +418,6 @@ s_noise: 1.0
# ETA (basically ancestralness).
eta: 1.0
# If the ETA calculation fails, it will retry with eta - eta_retry_increment until it either
# succeeds or eta becomes <= 0 (in which case ancestralness just gets disabled).
# In other words, you can set ETA as high as you want, set eta_retry_increment to something like 0.1 and
# it will just do whatever it takes to find an ETA that works.
eta_retry_increment: 0
# No effect unless both start and end are set. Will scale the eta value based on the
# percentage of sampling. In other words, eta*dyn_eta_start at the beginning,
# eta*dyn_eta_end at the end.
@@ -528,25 +435,15 @@ cfgpp: false
### Reversible Settings ###
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
# 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
pre_filter: null
@@ -629,99 +526,6 @@ 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 ###
@@ -734,13 +538,10 @@ weoon:
# 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:
# heunpp2: 3
# rk_dynamic: 4 - Can also be set to 0 to try to dynamically adjust the order based on calculated
# error from the last step (but that may not work well).
max_order: 999 # Varies based on sampler.
# Used for dpmpp_2m. One of midpoint, heun
@@ -769,50 +570,6 @@ dyn_deta_mode: "lerp"
</details>
### `OCS Param` and `OCS MultiParam`
Allows specifying parameter inputs that can't be expressed with YAML, such as custom noise.
#### Node Parameters
* `key`: Determines the parameter input type.
#### Input Parameters
* `value`: Input for the actual parameter - must match the specified input type or you will get an error when you evaluate the workflow.
* `params_opt`: Allows connecting another `OCS Param` or `OCS MultiParam` node to specify multiple parameters.
#### Text Parameters
Allows specifying extra parameters. You may use this block to rename a key, for example if your key type was `custom_noise` you could enter:
```yaml
rename: test
```
in the `OCS Param` node and:
```yaml
custom_noise: test
```
in the node that was connected to the params to have it use the custom noise specifically named `test`.
### `OCS MultiParam`
MultiParam is the same as Param except it has multiple optional inputs like `key_1`, `key_2`, etc.
#### Text Parameters
Same as `OCS Param` (see above), however if set you should use an object with a key corresponding to the index of the param. For example, if you wanted to rename `key_1` and `key_2` you would do something like:
```yaml
1:
rename: test1
2:
rename: test2
```
### `OCS SimpleRestartSchedule`
Generates a restart schedule.
@@ -877,6 +634,7 @@ For more tuneable parameters, see the `OCSNoise PerlinAdvanced` node.
* `res_height`: Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.
* `break_pattern`: 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.
***
### `OCSNoise PerlinAdvanced`
-4
View File
@@ -9,9 +9,5 @@ NODE_CLASS_MAPPINGS = {
"OCS MultiParam": nodes.MultiParamNode,
"OCS ModelSetMaxSigma": nodes.ModelSetMaxSigmaNode,
"OCS SimpleRestartSchedule": nodes.SimpleRestartSchedule,
"OCS ApplyFilterLatent": nodes.ApplyFilterLatent,
"OCS ApplyFilterImage": nodes.ApplyFilterImage,
"OCS ExpressionFilteredLatentOperation": nodes.ExpressionFilteredLatentOperationNode,
"OCS ExpressionFilteredModelPatch": nodes.ExpressionFilteredModelPatchNode,
} | custom_noise.NODE_CLASS_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS"]
+1 -3
View File
@@ -150,9 +150,7 @@ 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.
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
## Tensor 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.
+3 -6
View File
@@ -1,11 +1,8 @@
from . import noise_perlin
from . import nodes
NODE_CLASS_MAPPINGS = {
"OCSNoise PerlinSimple": nodes.PerlinSimpleNode,
"OCSNoise PerlinAdvanced": nodes.PerlinAdvancedNode,
"OCSNoise ImmiscibleReference": nodes.ImmiscibleReferenceNoiseNode,
"OCSNoise PerlinSimple": noise_perlin.PerlinSimpleNode,
"OCSNoise PerlinAdvanced": noise_perlin.PerlinAdvancedNode,
"OCSNoise to SONAR_CUSTOM_NOISE": nodes.ToSonarNode,
"OCSNoise Conditioning": nodes.NoiseConditioningNode,
"OCSNoise OverrideSamplerNoise": nodes.SamplerNodeConfigOverride,
"OCSNoise ExpressionFilteredNoise": nodes.ExpressionFilteredNoiseNode,
}
+5 -7
View File
@@ -1,10 +1,8 @@
import abc
from typing import Any, Callable
import torch
from ..external import IntegratedNode
from ..nodes import NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE
from typing import Callable, Any
from ..noise import scale_noise
@@ -107,7 +105,7 @@ class CustomNoiseChain:
return noise_sampler
class CustomNoiseNodeBase(metaclass=IntegratedNode):
class CustomNoiseNodeBase(abc.ABC):
DESCRIPTION = "An Overly Complicated Sampling custom noise item."
RETURN_TYPES = ("OCS_NOISE",)
OUTPUT_TOOLTIPS = ("A custom noise chain.",)
@@ -153,9 +151,9 @@ class CustomNoiseNodeBase(metaclass=IntegratedNode):
if include_chain:
result["optional"] |= {
"ocs_noise_opt": (
WILDCARD_NOISE,
"OCS_NOISE",
{
"tooltip": f"Optional input for more custom noise items.\n{NOISE_INPUT_TYPES_HINT}",
"tooltip": "Optional input for more custom noise items.",
},
),
}
+1 -1925
View File
File diff suppressed because it is too large Load Diff
-156
View File
@@ -1,156 +0,0 @@
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
File diff suppressed because it is too large Load Diff
+4 -8
View File
@@ -1,12 +1,8 @@
from . import expression, types, util
from .expression import Expression
from . import types, expression, handler, util, validation
try:
from . import handler, validation
from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext
from .validation import Arg, ValidateArg
except (ImportError, ModuleNotFoundError):
pass
from .expression import Expression
from .validation import Arg, ValidateArg
from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext
__all__ = (
"Arg",
+15 -37
View File
@@ -1,22 +1,18 @@
import operator
import re
import operator
from tqdm import tqdm
from .parser import ParseError, Parser, ParserSpec
from .parser import Parser, ParserSpec, ParseError
from .types import (
Empty,
ExpBase,
ExpBinOp,
ExpDict,
ExpFunAp,
ExpKV,
ExpMethodAp,
ExpOp,
ExpReturn,
ExpStatements,
ExpBinOp,
ExpSym,
ExpStatements,
ExpFunAp,
ExpTuple,
ExpDict,
ExpKV,
)
COMMA_PRECEDENCE = 2
@@ -38,11 +34,10 @@ class Expression:
| :> # Key value binop
| := # Assignment
| ; # Sequencing
| :: # Method call
| [?:] # Ternary
| \[ | ] # Index
| \.\.\. # Index ellipsis
| '[-\w.:=]+ # Symbol
| '[-\w.]+ # Symbol
| `?[a-z][\w.]*`? # Function/variable names
)
\s*
@@ -63,13 +58,10 @@ class Expression:
def eval(self, handlers, *args, **kwargs):
if self.expr != ExpOp("default"):
tqdm.write(f"* OCS: EVAL: {self.expr}")
print("\nEVAL", self.expr)
if not isinstance(self.expr, ExpBase):
return self.expr
try:
return self.expr.eval(handlers, *args, **kwargs)
except ExpReturn as expret:
return expret.args[0]
return self.expr.eval(handlers, *args, **kwargs)
def __len__(self):
return len(self.expr)
@@ -89,25 +81,20 @@ 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))
@@ -163,9 +150,9 @@ class ExprParserSpec(ParserSpec):
def split_funap_args(toks):
if not isinstance(toks, (list, tuple)):
return ExpTuple((toks,)), ExpDict()
return ExpTuple(t for t in toks if not isinstance(t, ExpKV)), ExpDict(
{str(t.k): t.v for t in toks if isinstance(t, ExpKV)}
)
return ExpTuple(t for t in toks if not isinstance(t, ExpKV)), ExpDict({
str(t.k): t.v for t in toks if isinstance(t, ExpKV)
})
@staticmethod
def null_constant(p, token, bp):
@@ -206,14 +193,6 @@ class ExprParserSpec(ParserSpec):
p.expect(")")
return make_funap(left, *cls.split_funap_args(args))
@classmethod
def left_methodcall(cls, p, token, left, bp):
methname = p.parse_until(31)
p.expect("(")
funap = cls.left_funcall(p, token=None, left=methname, bp=None)
funap.args = ExpTuple((Empty, *funap.args))
return ExpMethodAp(left, funap)
@staticmethod
def left_comma(p, token, left, bp):
if p.token == ")":
@@ -265,7 +244,6 @@ class ExprParserSpec(ParserSpec):
def populate(self):
self.add_left(31, self.left_funcall, ("(",))
self.add_left(31, self.left_index, ("[",))
self.add_left(31, self.left_methodcall, ("::",))
self.add_leftright(29, self.left_binop, ("**",))
self.add_null(27, self.null_prefixop, ("+", "-", "!"))
self.add_left(25, self.left_binop, ("*", "/"))
+19 -119
View File
@@ -1,11 +1,8 @@
import operator
import traceback
from tqdm import tqdm
from .types import Empty, ExpDict, ExpOp, ExpReturn, ExpTuple
from .validation import ValidateArg, Arg, ValidateError
from .types import Empty, ExpDict, ExpOp
from .util import torch
from .validation import Arg, ValidateArg, ValidateError
class HandlerError(Exception):
@@ -64,12 +61,9 @@ class BaseHandler:
def __call__(self, obj, *, getter):
try:
val = self.handle(obj, getter)
except ExpReturn:
raise
return self.validate_output(obj, val)
except Exception as exc:
tb = traceback.format_exc()
raise HandlerError(f'Error evaluating "{obj.name}": {exc!s}\n{tb}') from exc
return self.validate_output(obj, val)
raise HandlerError(f'Error evaluating "{obj.name}":\n {exc!r}') from exc
def safe_get(self, key, obj, getter=None, *, default=Empty):
str_key = isinstance(key, str)
@@ -184,7 +178,7 @@ class EqHandler(BinopLogicHandler):
return a1 == a2
class NeqHandler(EqHandler):
class NeqHandler(BinopLogicHandler):
def handle(self, *args, **kwargs):
return not super().handle(*args, **kwargs)
@@ -218,8 +212,6 @@ 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
@@ -262,22 +254,6 @@ class UnarySimpleMathHandler(SimpleMathHandler):
input_validators = (Arg.numeric("lhs"),)
class SimpleOpHandler(BaseHandler):
input_validators = (Arg.present("lhs"), Arg.present("rhs"))
def __init__(self, handler):
super().__init__()
self.handler = handler
def handle(self, obj, getter):
args = (
self.safe_get(idx, obj, getter=getter)
for idx in range(len(self.input_validators))
)
result = self.handler(*args)
return result
class IsSetHandler(BaseHandler):
input_validators = (Arg.string("name"),)
@@ -305,17 +281,9 @@ class GetHandler(BaseHandler):
class S_Handler(BaseHandler):
input_validators = (
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),
Arg.integer("start", None),
Arg.integer("end", None),
Arg.integer("step", None),
)
def handle(self, obj, getter):
@@ -350,13 +318,6 @@ class MaxHandler(MinHandler):
return max(*self.safe_get("values", obj, getter))
class SumHandler(BaseHandler):
input_validators = (Arg.numeric_sequence("values"),)
def handle(self, obj, getter):
return sum(tuple(self.safe_get("values", obj, getter)))
class UnsafeCallHandler(BaseHandler):
input_validators = (Arg.present("__callable"),)
@@ -367,7 +328,7 @@ class UnsafeCallHandler(BaseHandler):
)
fun = self.safe_get("__callable", obj, getter)
if not callable(fun):
raise TypeError("Cannot call supplied value: not a callable")
raise ValueError("Cannot call supplied value: not a callable")
args = (self.safe_get(idx, obj, getter) for idx in range(1, len(obj.args)))
kwargs = {k: self.safe_get(k, obj, getter) for k in obj.kwargs}
return fun(*args, **kwargs)
@@ -394,53 +355,6 @@ class SetVarHandler(BaseHandler):
return val
class ReturnHandler(BaseHandler):
input_validators = (Arg.present("expression"),)
def handle(self, obj, getter):
raise ExpReturn(self.safe_get("expression", obj, getter))
class PrintHandler(BaseHandler):
input_validators = (Arg.present("lhs"),)
def handle(self, obj, getter):
lhs = self.safe_get("lhs", obj, getter)
tqdm.write(f"[OCS expr_print]: {lhs!s}")
class MapHandler(BaseHandler):
input_validators = (
Arg.sequence("items"),
Arg.string("key", default="item"),
Arg.present("expression"),
Arg.present("check_expression", default=None),
)
class _MapEmpty:
pass
def handle(self, obj, getter):
items = self.safe_get("items", obj, getter)
key = self.safe_get("key", obj, getter)
have_check_expr = None
result = []
for item in items:
getter.ctx.set_var(key, item)
if have_check_expr in {True, None}:
checked = self.safe_get(
"check_expression",
obj,
getter,
default=self._MapEmpty,
)
have_check_expr = checked is not self._MapEmpty
if have_check_expr and not bool(checked):
continue
result.append(self.safe_get("expression", obj, getter))
return ExpTuple(result)
LOGIC_HANDLERS = {
"||": OrHandler(),
"&&": AndHandler(),
@@ -461,26 +375,21 @@ for k, alias in (
MATH_HANDLERS = {
"*": SimpleMathHandler(operator.mul),
"**": SimpleMathHandler(operator.pow),
"+": SimpleMathHandler(operator.add),
"-": MinusHandler(),
"*": SimpleMathHandler(operator.mul),
"/": SimpleMathHandler(operator.truediv),
"//": SimpleMathHandler(operator.floordiv),
"**": SimpleMathHandler(operator.pow),
"mod": SimpleMathHandler(operator.mod),
"neg": UnarySimpleMathHandler(operator.neg),
"between": BetweenHandler(),
"<": RelComparisonHandler(operator.lt),
"<=": RelComparisonHandler(operator.le),
">": RelComparisonHandler(operator.gt),
">=": RelComparisonHandler(operator.ge),
"abs": SimpleMathHandler(operator.abs),
"between": BetweenHandler(),
"bool": UnarySimpleMathHandler(handler=bool),
"float": UnarySimpleMathHandler(handler=float),
"int": UnarySimpleMathHandler(handler=int),
"max": MaxHandler(),
"min": MinHandler(),
"mod": SimpleMathHandler(operator.mod),
"neg": UnarySimpleMathHandler(operator.neg),
"sum": SumHandler(),
"max": MaxHandler(),
}
for k, alias in (
("+", "add"),
@@ -493,23 +402,14 @@ for k, alias in (
MATH_HANDLERS[alias] = MATH_HANDLERS[k]
MISC_HANDLERS = {
"and": SimpleOpHandler(operator.and_),
"comment": CommentHandler(),
"concat": SimpleOpHandler(operator.concat),
"contains": SimpleOpHandler(operator.contains),
"dict": DictHandler(),
"is_set": IsSetHandler(),
"get": GetHandler(),
"index": IndexHandler(),
"is_set": IsSetHandler(),
"map": MapHandler(),
"op_or": SimpleOpHandler(operator.or_),
"op_and": SimpleOpHandler(operator.and_),
"op_xor": SimpleOpHandler(operator.xor),
"print": PrintHandler(),
"return": ReturnHandler(),
"s_": S_Handler(),
"set_var": SetVarHandler(),
"unsafe_call": UnsafeCallHandler(),
"dict": DictHandler(),
"comment": CommentHandler(),
"set_var": SetVarHandler(),
}
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
+30 -98
View File
@@ -3,10 +3,6 @@ class Empty:
return False
class ExpReturn(Exception):
pass
class ExpBase:
def __bool__(self):
return True
@@ -45,10 +41,8 @@ class ExpSym(str, ExpBase):
class ExpTuple(tuple, ExpBase):
__slots__ = ()
def clone(self, **kwargs):
return self.__class__(
v.clone(**kwargs) if isinstance(v, ExpBase) else v for v in self
)
def clone(self):
return self.__class__(v.clone() if isinstance(ExpBase) else v for v in self)
def get_eval(self, k, handlers, *args, default=None, **kwargs):
val = super().__getitem__(k)
@@ -83,10 +77,11 @@ class ExpKV(ExpBase):
class ExpDict(dict, ExpBase):
__slots__ = ()
def clone(self, **kwargs):
return self.__class__(
v.clone(**kwargs) if isinstance(v, ExpBase) else v for v in self
)
def clone(self):
return self.__class__(v.clone() if isinstance(ExpBase) else v for v in self)
def pop(self, *args, **kwargs):
raise NotImplementedError
def get_eval(self, k, handlers, *args, default=Empty, **kwargs):
val = super().get(k, default)
@@ -111,17 +106,12 @@ class ExpDict(dict, ExpBase):
for k, v in self.items()
}
# Can't remember if there was a compelling reason ExpDict can't be mutable but
# it breaks deep copy stuff.
#
# def pop(self, *args, **kwargs):
# raise NotImplementedError
# popitem = pop
# update = pop
# clear = pop
# __delitem__ = pop
# __setitem__ = pop
# __ior__ = pop
popitem = pop
update = pop
clear = pop
__delitem__ = pop
__setitem__ = pop
__ior__ = pop
class ExpStatements(ExpBase):
@@ -145,81 +135,26 @@ class ExpStatements(ExpBase):
class ExprGetter:
def __init__(self, obj, ctx, args, kwargs, *, prepend_args=()):
def __init__(self, obj, ctx, *args, **kwargs):
self.obj = obj
self.ctx = ctx
self.args = args
self.prepend_args = prepend_args
self.kwargs = kwargs
def __call__(self, k, *, default=Empty):
obj = self.obj
if isinstance(k, str):
result = obj.kwargs.get_eval(
k, self.ctx, *self.args, default=default, **self.kwargs
)
elif isinstance(k, int):
pa = self.prepend_args
pa_len = len(pa)
result = (
pa[k]
if k < pa_len
else obj.args.get_eval(k, self.ctx, *self.args, **self.kwargs)
)
result = (
obj.kwargs.get_eval(k, self.ctx, *self.args, default=default, **self.kwargs)
if isinstance(k, str)
else obj.args.get_eval(k, self.ctx, *self.args, **self.kwargs)
)
if result is Empty:
raise KeyError(f"Unknown key {k!r}")
return result
class ExpMethodAp(ExpBase):
__slots__ = ("funap", "object_expression")
def __init__(self, object_expression, funap):
super().__init__()
self.object_expression = object_expression
self.funap = funap
def eval(self, handlers, *args, **kwargs):
object_value = self.object_expression.eval(handlers, *args, **kwargs)
type_name = type(object_value).__name__
handler_key = f"{type_name}::{self.funap.name}"
handler = handlers.get_handler(handler_key)
if handler is Empty:
raise KeyError(f"No handler for method call op: {handler_key!r}")
return handler(
self,
getter=ExprGetter(
self.funap, handlers, args, kwargs, prepend_args=(object_value,)
),
**kwargs,
)
def clone(self, **kwargs):
return self.__class__(
object_expression=self.object_expression.clone(**kwargs),
funap=self.funap.clone(**kwargs),
)
__copy__ = clone
def __getattr__(self, k):
if k == "name":
return f"method::{self.funap.name}"
if k == "args":
return self.funap.args
if k == "kwargs":
return self.funap.kwargs
# This doesn't play well with deep copy.
# if hasattr(self.funap, k):
# return getattr(self.funap, k)
raise AttributeError(f"Can't get attribute {k}")
def __repr__(self):
return f"<METHAP:{self.object_expression}::{self.funap}>"
class ExpFunAp(ExpBase):
__slots__ = ("args", "kwargs", "name")
__slots__ = ("name", "args", "kwargs")
def __init__(self, name, args=None, kwargs=None):
self.name = name
@@ -230,15 +165,13 @@ class ExpFunAp(ExpBase):
handler = handlers.get_handler(self.name)
if handler is Empty:
raise KeyError(f"No handler for op: {self.name!r}")
return handler(self, getter=ExprGetter(self, handlers, args, kwargs), **kwargs)
def clone(self, **kwargs):
return self.__class__(
self.name,
self.args.clone(**kwargs),
self.kwargs.clone(**kwargs),
return handler(
self, getter=ExprGetter(self, handlers, *args, **kwargs), **kwargs
)
def clone(self):
return self.__class__(self.name, self.args.clone(), self.kwargs.clone())
def pretty_string(self, depth=0):
pad = " " * (depth + 1) * 2
kwargs_str = f", {self.kwargs.pretty_string(depth + 1)}" if self.kwargs else ""
@@ -269,13 +202,12 @@ class ExpBoundFunAp(ExpFunAp):
__all__ = (
"ExpBase",
"ExpBinOp",
"ExpBoundFunAp",
"ExpDict",
"ExpFunAp",
"ExpKV",
"ExpMethodAp",
"ExpOp",
"ExpBinOp",
"ExpSym",
"ExpTuple",
"ExpKV",
"ExpDict",
"ExpFunAp",
"ExpBoundFunAp",
)
+21 -101
View File
@@ -1,13 +1,12 @@
import contextlib
import functools
from ..latent import ImageBatch
from .types import Empty
from .util import torch
from .types import Empty
class Arg:
__slots__ = ("default", "name", "validator")
__slots__ = ("name", "default", "validator")
def __init__(self, name, default=Empty, *, validator=None):
self.name = name
@@ -53,27 +52,6 @@ class Arg:
name, default=default, validator=ValidateArg.validate_numscalar_sequence
)
@classmethod
def numeric_sequence(cls, name, default=Empty):
return cls(
name, default=default, validator=ValidateArg.validate_numeric_sequence
)
@classmethod
def numscalar_sequence_or_single(cls, name, default=Empty):
return cls.one_of(
name,
(
ValidateArg.validate_numscalar_sequence,
ValidateArg.validate_numeric_scalar,
),
default=default,
)
@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(
@@ -84,16 +62,6 @@ 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)
@@ -103,8 +71,8 @@ class Arg:
return cls(name, default=default, validator=ValidateArg.validate_boolean)
@classmethod
def present(cls, name, default=Empty):
return cls(name, default=default, validator=ValidateArg.validate_passthrough)
def present(cls, name):
return cls(name, validator=ValidateArg.validate_passthrough)
@classmethod
def one_of(cls, name, validators, *, default=Empty):
@@ -126,14 +94,11 @@ class ValidateError(Exception):
class ValidateArg:
__slots__ = ("groupfun", "kwargs", "kwargslist", "valfuns")
__slots__ = ("valfuns", "groupfun", "kwargs", "kwargslist")
def __init__(self, name, *args, kwargslist=(), group=all, **kwargs):
if not isinstance(name, (list, tuple)):
name = (name,)
args = ((args,),)
kwargslist = kwargs
kwargs = {}
return self.__init__((name,), (args,), group=group, kwargslist=kwargs)
self.valfuns = (getattr(self, f"validate_{n}", None) for n in name)
if not all(self.valfuns):
raise ValueError("Unknown validator")
@@ -166,22 +131,6 @@ 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):
@@ -211,29 +160,7 @@ 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 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}"
)
raise ValidateError(f"Item validation failed for in sequence: {exc}")
@classmethod
def validate_numscalar_sequence(cls, idx, val):
@@ -241,15 +168,20 @@ class ValidateArg:
idx, val, item_validator=cls.validate_numeric_scalar
)
@classmethod
def validate_numeric_sequence(cls, idx, val):
return cls.validate_sequence(idx, val, item_validator=cls.validate_numeric)
@classmethod
def validate_tensor_slice(cls, idx, val):
return cls.validate_sequence(
idx, val, item_validator=cls.validate_tensor_slice_item
)
# @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_string(cls, idx, val):
@@ -257,24 +189,12 @@ class ValidateArg:
raise ValidateError(f"Expected string argument at {idx}, got {type(val)}")
return val
@classmethod
def validate_dict(cls, idx, val):
if not isinstance(val, dict):
raise ValidateError(f"Expected dict argument at {idx}, got {type(val)}")
return val
@classmethod
def validate_boolean(cls, idx, val):
if val is not True and val is not False:
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
+420 -757
View File
File diff suppressed because it is too large Load Diff
+13 -112
View File
@@ -1,120 +1,21 @@
import contextlib
import importlib
import sys
from functools import partial
from types import ModuleType
from typing import Callable, NamedTuple
MODULES = {}
class Integrations:
class Integration(NamedTuple):
key: str
module_name: str
handler: Callable | None = None
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
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["sonar"] = importlib.import_module("custom_nodes.ComfyUI-sonar").py
with contextlib.suppress(ImportError, NotImplementedError):
MODULES["nnlatentupscale"] = importlib.import_module(
"custom_nodes.ComfyUi_NNLatentUpscale"
)
__all__ = ("MODULES",)
+116 -130
View File
@@ -10,33 +10,23 @@ from .utils import fallback
OD = collections.OrderedDict
BLENDING_MODES = {
"lerp": torch.lerp,
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 | {
"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
@@ -44,9 +34,8 @@ FILTER_HANDLERS = expr.HandlerContext(
class FilterRefs:
def __init__(self, kvs=None, *, ctx=None):
def __init__(self, kvs=None):
self.kvs = fallback(kvs, {})
self.ctx = fallback(ctx, {})
def get(self, k, default=None):
return self.kvs.get(k, default)
@@ -58,14 +47,13 @@ class FilterRefs:
self.kvs[k] = v
def clone(self):
return self.__class__(self.kvs.copy(), ctx=self.ctx)
return self.__class__(self.kvs.copy())
def __or__(self, other):
return self.__class__(self.kvs | other.kvs, ctx=self.ctx | other.ctx)
return self.__class__(self.kvs | other.kvs)
def __ior__(self, other):
self.kvs |= other.kvs
self.ctx |= other.ctx
return self
def __delitem__(self, k):
@@ -83,34 +71,26 @@ class FilterRefs:
def __iter__(self):
return self.kvs.__iter__()
def items(self):
return self.kvs.items()
@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,
},
ctx={"ss": ss, "model": ss.model},
)
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,
})
if have_current and len(ss.hist) > 0:
fr |= cls.from_mr(ss.hcur)
fr["d"] = ss.d
@@ -246,12 +226,6 @@ 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"<Filter({self.name}): {prettyvals}>"
class SimpleFilter(Filter):
name = "simple"
@@ -419,88 +393,100 @@ 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,
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,
}
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
)
if EXT_SONAR:
class SonarPowerFilter(Filter):
name = "sonar_power_filter"
default_options = Filter.default_options
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)
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
)
)
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"]
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)
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)
FILTER |= {"sonar_power_filter": SonarPowerFilter}
def make_filter(args):
+107 -484
View File
@@ -1,30 +1,16 @@
from typing import Any, NamedTuple, Self
import folder_paths
import latent_preview
import numpy as np
import torch
import torch.nn.functional as F
from comfy import latent_formats
import folder_paths
import latent_preview
from comfy.taesd.taesd import TAESD
from comfy.utils import bislerp
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 = (
@@ -39,33 +25,32 @@ 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.
# 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(
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",
)
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
def on_abs_stacked(tensor_list, f, *args: list, **kwargs: dict):
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):
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
@@ -117,12 +102,7 @@ def contrast_adaptive_sharpening(
div = torch.reciprocal(1 + 4 * w)
output = ((b + d + f + h) * w + e) * div
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)
return output.real.clamp(x.min(), x.max())
class ImageBatch(tuple):
@@ -165,11 +145,24 @@ 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])
@@ -197,450 +190,80 @@ 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)
bleh_scale_samples = None
UPSCALE_METHODS = ("bicubic", "bislerp", "bilinear", "nearest-exact", "area")
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)
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 "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 get_noise_sampler(noise_type, x, *_args: list, **_kwargs: dict): # noqa: F811
if noise_type != "gaussian":
raise ValueError("Only gaussian noise supported unless you have ComfyUI-sonar")
return lambda _s, _sn: torch.randn_like(x)
if "nnlatentupscale" in EXT:
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
)
/ 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.rgb_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.rgb_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)
def randomized_svd(
m: torch.Tensor,
*,
rank: int | None = None,
n_iter: int = 6,
ortho_interval: int = 3,
oversample: int = 10,
noise_sampler: Callable | None = None,
y: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
n, c = m.shape[-2], m.shape[-1]
if rank is None:
rank = n
if y is None:
if noise_sampler is None:
noise_sampler = torch.randn
k = min(rank + oversample, n, c)
y_shape = (*m.shape[:-2], c, k)
y = noise_sampler(y_shape, device=m.device, dtype=m.dtype)
elif y.shape == m.shape:
y = y.mT
elif y.shape != m.mT.shape:
raise ValueError("Bad initial y shape")
y = m @ y
ortho = False
for idx in range(n_iter):
y = m @ (m.mT @ y)
ortho = ortho_interval > 0 and (idx % ortho_interval) == 0 and n_iter - idx != 2
if ortho:
y = torch.linalg.qr(y)[0]
q = y if ortho else torch.linalg.qr(y)[0]
u, s, vh = torch.linalg.svd(q.mT @ m, full_matrices=False)
u = q @ u
if rank < n:
return u[..., :rank], s[..., :rank], vh[..., :rank, :]
return u, s, vh
class DimCorrelationOrder(NamedTuple):
perm: torch.Tensor
dim: int
leave: bool = False
def reorder(
self,
x: torch.Tensor,
def scale_nnlatentupscale(
mode,
latent,
scale=2.0,
*,
invert: bool = False,
) -> torch.Tensor:
dim, perm = self.dim, self.perm
if invert:
perm = perm.argsort(dim=-1)
if dim < 0:
dim = x.ndim + self.dim
shape = [1] * x.ndim
if dim != 0:
shape[0] = x.shape[0]
shape[dim] = x.shape[dim]
perm = perm.view(*shape).expand_as(x)
return x.gather(dim=dim, index=perm)
class DimCorrelationConfig(NamedTuple):
dim: int = 1
flip: bool = False
cross: bool = False
leave: bool = False
preserve_first: bool = False
center_strength: float = 1.0
center_dim: int = -1
# None - disabled, otherwise controls whether abs occurs before or after centering.
abs_before: bool | None = None
# 0 - disabled, positive value - enabled, negative value - enabled with sign flipped.
fix_sign: int = 1
align_to_peak: bool = False
expansion_factor: int = 1 # NYI
# Only applies to cross mode. One of: svd, randomized_svd
decomp_mode: str = "svd"
low_rank: int = 0
low_rank_niter: int = 6
work_dtype: torch.dtype | None = None
pc: int = 0
@classmethod
def build(cls, **kwargs: Any) -> Self:
wd = kwargs.get("work_dtype")
if isinstance(wd, str):
dtype_map = {
"float64": torch.float64,
"float32": torch.float32,
"float16": torch.float16,
"bfloat16": torch.bfloat16,
}
wd = dtype_map.get(wd)
if wd is None:
raise ValueError("Bad dtype")
kwargs["work_dtype"] = wd
if kwargs.get("decomp_mode") not in {None, "svd", "randomized_svd"}:
raise ValueError(
"Bad decomp mode, must be unset or one of: svd, randomized_svd",
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
)
fs = frozenset(cls._fields)
kwargs = {k: v for k, v in kwargs.items() if k in fs}
return cls(**kwargs)
@classmethod
def from_str(cls, s: str) -> Self:
parts = tuple(p.strip() for p in s.split(":", 3))
plen = len(parts)
if plen > 2:
raise ValueError("Dim correlations only support up to two parts.")
dim = int(parts[0])
result = cls(dim=dim)
if plen < 2 or not parts[1]:
return result
p1 = parts[1]
p1len = len(p1)
offs = 0
while offs < p1len:
pflag = p1[offs]
offs += 1
if pflag == "f":
result = result._replace(flip=True)
elif pflag == "x":
result = result._replace(cross=True)
elif pflag == "l":
result = result._replace(leave=True)
elif pflag == "u":
result = result._replace(center_strength=0.0)
elif pflag == "c":
result = result._replace(center_dim=-2)
elif pflag in "aA":
result = result._replace(abs_before=pflag == "a")
elif pflag == "s":
result = result._replace(fix_sign=0)
elif pflag == "S":
result = result._replace(fix_sign=-1)
elif pflag == "p":
result = result._replace(align_to_peak=True)
elif pflag == "i":
result = result._replace(preserve_first=True)
else:
offs -= 1
break
p1 = p1[offs:].strip()
pc = int(p1) if p1 else 0
return result._replace(pc=pc)
@classmethod
def _fix_sign_ambiguity(
cls,
x: torch.Tensor,
*,
ref: torch.Tensor | None = None,
dim: int = -1,
in_place: bool = True,
neg: bool = False,
) -> torch.Tensor:
if ref is not None:
x = cls._fix_sign_ambiguity(x, dim=dim, in_place=in_place)
else:
ref = x
signs = ref.gather(dim, ref.abs().argmax(dim=dim, keepdim=True)).sign_()
signs = signs.masked_fill_(signs == 0, 1.0)
if neg:
signs = signs.neg_()
return x.mul_(signs) if in_place else x * signs
@staticmethod
def _preserve_first_index(perm: torch.Tensor) -> torch.Tensor:
c = perm.shape[-1]
shift = (perm == 0).to(dtype=torch.int64).argmax(dim=-1, keepdim=True)
shift = torch.arange(c, device=perm.device, dtype=shift.dtype) + shift
shift %= c
return perm.gather(dim=-1, index=shift)
@staticmethod
def _align_ref(
*,
x: torch.Tensor,
ref: torch.Tensor,
skip_dim: int,
) -> torch.Tensor:
if skip_dim < 0:
skip_dim = x.ndim + skip_dim
reps = tuple(
None if szr == 0 or szx % szr != 0 else (d, szx // szr)
for d, (szx, szr) in enumerate(zip(x.shape, ref.shape, strict=True))
if d != skip_dim and szx != szr
/ scale_factor
)
if not all(reps):
raise ValueError("Bad shape")
for d, r in reps:
ref = ref.repeat_interleave(r, d)
return ref
def _get_correlation_order(
self,
cov: torch.Tensor,
*,
pc_idx: int = 0,
cross_mode: bool = False,
**kwargs: Any,
) -> torch.Tensor:
if cross_mode:
if self.decomp_mode == "randomized_svd":
pcs = randomized_svd(cov, n_iter=self.low_rank_niter, **kwargs)[0]
elif self.low_rank < 1:
pcs = torch.linalg.svd(cov, full_matrices=False).U
else:
pcs = torch.svd_lowrank(
cov,
q=self.low_rank,
niter=self.low_rank_niter,
)[0]
n_pcs = pcs.shape[-1]
if pc_idx < 0:
pc_idx = n_pcs + pc_idx
else:
pcs = torch.linalg.eigh(cov).eigenvectors
n_pcs = pcs.shape[-1]
pc_idx = n_pcs - pc_idx - 1 if pc_idx >= 0 else pc_idx + 1
pc_idx = max(0, min(n_pcs - 1, pc_idx))
return pcs[..., pc_idx]
def _preprocess(
self,
x: torch.Tensor,
*,
dim: int,
allow_post_abs: bool = True,
) -> torch.Tensor:
x_flat = (x.unsqueeze(0) if dim == 0 else x.movedim(dim, 1)).flatten(
start_dim=2
)
if self.abs_before is True:
x_flat = x_flat.abs()
if self.center_strength == 0:
return x_flat
xm = x_flat.mean(dim=self.center_dim, keepdim=True)
if self.center_strength != 1:
xm *= self.center_strength
x_flat = x_flat.sub_(xm) if self.abs_before else x_flat - xm
if allow_post_abs and self.abs_before is False:
x_flat = x_flat.abs_()
return x_flat
def get_correlation_order(
self,
x: torch.Tensor,
*,
ref: torch.Tensor | None = None,
**kwargs: Any,
) -> DimCorrelationOrder:
if self.work_dtype is not None:
if x.dtype != self.work_dtype:
x = x.to(dtype=self.work_dtype)
if ref is not None and ref.dtype != self.work_dtype:
ref = ref.to(dtype=self.work_dtype)
dim, pc_idx = self.dim, self.pc
if dim < 0:
dim = x.ndim + dim
allow_post_abs = self.abs_before is not False or not self.align_to_peak
x_flat = self._preprocess(
x,
dim=dim,
allow_post_abs=ref is not None or allow_post_abs,
)
if ref is None or not self.cross:
# if ref is None:
y_flat = x_flat
if not allow_post_abs:
sign_ref = y_flat.clone()
y_flat = y_flat.abs_()
else:
sign_ref = y_flat if self.align_to_peak else None
else:
if ref.shape != x.shape:
ref = self._align_ref(x=x, ref=ref, skip_dim=dim)
y_flat = self._preprocess(ref, dim=dim, allow_post_abs=allow_post_abs)
if not allow_post_abs:
sign_ref = y_flat.clone()
y_flat = y_flat.abs_()
else:
sign_ref = y_flat if self.align_to_peak else None
pc = self._get_correlation_order(
x_flat @ y_flat.mT,
pc_idx=pc_idx,
cross_mode=ref is not None,
**kwargs,
)
if self.fix_sign:
pc = self._fix_sign_ambiguity(
pc,
neg=self.fix_sign < 0,
ref=None
if sign_ref is None
else sign_ref.flatten(start_dim=0 if dim == 0 else 1),
)
perm = pc.argsort(dim=-1, descending=self.flip)
if self.preserve_first:
perm = self._preserve_first_index(perm)
return DimCorrelationOrder(dim=dim, perm=perm, leave=self.leave)
del model
return result
+76 -169
View File
@@ -1,3 +1,5 @@
from collections import namedtuple
import torch
import comfy
@@ -6,7 +8,6 @@ from comfy.k_diffusion.sampling import to_d
from . import filtering
from .utils import fallback
from .latent import OCSLatentFormat
class History:
@@ -42,14 +43,12 @@ 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:
@@ -73,33 +72,6 @@ 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()
@@ -122,44 +94,27 @@ class ModelResult:
setattr(obj, k, val)
return obj
def get_error(self, other, *, override=None, alt_cfgpp_scale=0, cfgpp=False):
slf = fallback(override, self)
first, second = (other, slf) if other.sigma > slf.sigma else (slf, other)
if first.sigma == second.sigma:
return 0.0
d = first.to_d(alt_cfgpp_scale=alt_cfgpp_scale, cfgpp=cfgpp)
d_pred = second.to_d(
x=first.x + d * (second.sigma - first.sigma),
alt_cfgpp_scale=alt_cfgpp_scale,
cfgpp=cfgpp,
)
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 OCSModel:
class ModelCallCache:
def __init__(
self,
model,
x: torch.Tensor,
s_in: torch.Tensor,
extra_args: dict,
x,
s_in,
extra_args,
*,
cache: dict | None = None,
filter: dict | None = None,
cfg1_uncond_optimization: bool = False,
cfg_scale_override: int | float | None = None,
) -> None:
cache=None,
filter=None,
):
self.cache = ModelCallCacheConfig(**fallback(cache, {}))
filtargs = fallback(filter, {}).copy()
self.filters = {}
for key in (
"input",
"denoised",
"jdenoised",
"cond",
"uncond",
"postcfg",
"precfg",
):
for key in ("input", "denoised", "jdenoised", "cond", "uncond", "x"):
filt = filtargs.pop(key, None)
if filt is None:
continue
@@ -167,35 +122,25 @@ class OCSModel:
self.model = model
self.s_in = s_in
self.extra_args = extra_args
self.cfg1_uncond_optimization = cfg1_uncond_optimization
self.cfg_scale_override = cfg_scale_override
self.model_sampling = model.inner_model.inner_model.model_sampling
self.is_rectified_flow = isinstance(
self.model_sampling, comfy.model_sampling.CONST
)
self.latent_format = OCSLatentFormat(
x.device, model.inner_model.inner_model.latent_format
)
if self.cache.size < 1:
return
self.reset_cache()
def maybe_filter(
self, name: str, latent: torch.Tensor, *args: list, **kwargs: dict
) -> torch.Tensor:
def maybe_filter(self, name, latent, *args, **kwargs):
filt = self.filters.get(name)
if filt is None:
return latent
return filt.apply(latent, *args, **kwargs)
def filter_result(
self, result: ModelResult, *args: list, **kwargs: dict
) -> ModelResult:
def filter_result(self, result, *args, **kwargs):
if not self.filters:
return result
result = result.clone()
for key in ("denoised", "cond", "uncond", "jdenoised"):
for key in ("denoised", "cond", "uncond", "jdenoised", "x"):
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
@@ -203,123 +148,87 @@ class OCSModel:
return result
@staticmethod
def _fr_add_mr(fr: filtering.FilterRefs, mr: ModelResult) -> filtering.FilterRefs:
def _fr_add_mr(fr, mr):
frmr = filtering.FilterRefs.from_mr(mr)
fr.kvs |= {f"{k}_curr": v for k, v in frmr.kvs.items()}
return fr
def call_model(
self, x: torch.Tensor, sigma: torch.Tensor, **kwargs: dict
) -> torch.Tensor:
def reset_cache(self):
size = self.cache.size
self.slot = [None] * size
self.slot_use = [self.cache.max_use] * size
def get(self, idx, *, jvp=False):
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, mr):
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, sigma, **kwargs):
return self.model(x, sigma * self.s_in, **self.extra_args | kwargs)
# @property
# def model_sampling(self):
# return self.model.inner_model.inner_model.model_sampling
@property
def inner_cfg_scale(self) -> None | int | float:
maybe_cfg_scale = getattr(self.model.inner_model, "cfg", None)
return maybe_cfg_scale if isinstance(maybe_cfg_scale, (int, float)) else None
def set_inner_cfg_scale(self, scale: None | int | float) -> None | int | float:
eff_scale = self.cfg_scale_override
if scale is not None:
eff_scale = None if scale < 0 else scale
if eff_scale is None or eff_scale < 0:
return None
curr_cfg_scale = self.inner_cfg_scale
if curr_cfg_scale is None:
return None
self.model.inner_model.cfg = eff_scale
return curr_cfg_scale
def model_sampling(self):
return self.model.inner_model.inner_model.model_sampling
def __call__(
self,
x: torch.Tensor,
sigma: torch.Tensor,
x,
sigma,
*,
call_index: int = 0,
call_index=0,
ss,
s_in=None,
tangents=None,
require_uncond: bool = False,
cfg_scale_override: None | int = None,
return_cached=False,
**kwargs,
) -> ModelResult:
):
filter_refs = ss.refs | filtering.FilterRefs({
"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, True) if return_cached else 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, have_uncond
nonlocal denoised_cond, denoised_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 result
return args["denoised"]
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)
})
extra_args = self.extra_args | {
"model_options": comfy.model_patcher.set_model_options_post_cfg_function(
model_options, postcfg, disable_cfg1_optimization=True
)
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)
model_options = comfy.model_patcher.set_model_options_post_cfg_function(
model_options,
postcfg,
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):
@@ -327,31 +236,29 @@ class OCSModel:
if tangents is None:
denoised = call_model(x, sigma, **kwargs)
self.set_inner_cfg_scale(orig_cfg_scale)
mr = ModelResult(
call_index,
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
return (mr, False) if return_cached else mr
denoised, denoised_prime = torch.func.jvp(call_model, (x, sigma), tangents)
self.set_inner_cfg_scale(orig_cfg_scale)
mr = ModelResult(
call_index,
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
return (mr, False) if return_cached else mr
+51 -571
View File
@@ -1,77 +1,21 @@
from __future__ import annotations
import yaml
import comfy
import torch
import yaml
from tqdm import tqdm
from .external import MODULES, IntegratedNode
from .filtering import Filter, FilterRefs, make_filter
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 .substep_sampling import ParamGroup, StepSamplerChain, StepSamplerGroups
from .restart import Restart
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"
DEFAULT_YAML_PARAMS = """\
# JSON or YAML parameters
s_noise: 1.0
eta: 1.0
"""
class SamplerNode(metaclass=IntegratedNode):
class SamplerNode:
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."
@@ -88,8 +32,7 @@ class SamplerNode(metaclass=IntegratedNode):
"groups": (
"OCS_GROUPS",
{
"tooltip": "Connect OCS substep groups here which are output from the OCS Group node.",
"forceInput": True,
"tooltip": "Connect OCS substep groups here which are output from the OCS Group node."
},
),
},
@@ -98,14 +41,12 @@ class SamplerNode(metaclass=IntegratedNode):
"OCS_PARAMS",
{
"tooltip": "Optionally connect parameters like custom noise here. Output from the OCS Param or OCS MultiParam nodes.",
"forceInput": True,
},
),
"parameters": (
"STRING",
{
"default": "",
"placeholder": DEFAULT_YAML_PARAMS,
"default": 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.",
@@ -121,7 +62,6 @@ class SamplerNode(metaclass=IntegratedNode):
params_opt=None,
parameters="",
):
MODULES.initialize()
options = {}
parameters = parameters.strip()
if parameters:
@@ -140,7 +80,7 @@ class SamplerNode(metaclass=IntegratedNode):
)
class GroupNode(metaclass=IntegratedNode):
class GroupNode:
RETURN_TYPES = ("OCS_GROUPS",)
CATEGORY = "sampling/custom_sampling/OCS"
DESCRIPTION = "Over Complicated Sampling group definition node."
@@ -190,7 +130,6 @@ class GroupNode(metaclass=IntegratedNode):
"OCS_SUBSTEPS",
{
"tooltip": "Connect output from an OCS Substeps node here.",
"forceInput": True,
},
),
},
@@ -199,21 +138,18 @@ class GroupNode(metaclass=IntegratedNode):
"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": "",
"placeholder": DEFAULT_YAML_PARAMS,
"default": 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.",
@@ -234,7 +170,6 @@ class GroupNode(metaclass=IntegratedNode):
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
@@ -255,7 +190,7 @@ class GroupNode(metaclass=IntegratedNode):
return (group,)
class SubstepsNode(metaclass=IntegratedNode):
class SubstepsNode:
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."
@@ -290,21 +225,18 @@ class SubstepsNode(metaclass=IntegratedNode):
"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": "",
"placeholder": f"{DEFAULT_YAML_PARAMS}s_noise: 1.0\neta: 1.0\n",
"default": 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.",
@@ -321,7 +253,6 @@ class SubstepsNode(metaclass=IntegratedNode):
params_opt=None,
**kwargs,
):
MODULES.initialize()
if substeps_opt is not None:
chain = substeps_opt.clone()
else:
@@ -339,23 +270,30 @@ class SubstepsNode(metaclass=IntegratedNode):
return (chain,)
class ParamNode(metaclass=IntegratedNode):
class Wildcard(str):
__slots__ = ()
def __ne__(self, _unused):
return False
class ParamNode:
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_TOOLTIPS = (
OUTPUT_TYPES = (
"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"
OCS_PARAM_INPUT_TYPES = {
WC = Wildcard("*")
OCS_PARAM_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,
"vae": lambda _v: True,
"upscale_model": lambda _v: True,
"SAMPLER": lambda _v: True,
}
@classmethod
@@ -363,16 +301,15 @@ class ParamNode(metaclass=IntegratedNode):
return {
"required": {
"key": (
tuple(cls.OCS_PARAM_INPUT_TYPES.keys()),
tuple(cls.OCS_PARAM_TYPES.keys()),
{
"tooltip": "Used to set the type of custom parameter.",
},
),
"value": (
WILDCARD_PARAM,
cls.WC,
{
"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,
"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.",
},
),
},
@@ -381,47 +318,28 @@ class ParamNode(metaclass=IntegratedNode):
"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": "",
"placeholder": "# Additional YAML or JSON parameters",
"default": "# Additional YAML or JSON parameters\n",
"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.",
},
),
},
}
@classmethod
def get_renamed_key(cls, key, params):
rename = params.get("rename")
if rename is None:
return key
if not isinstance(rename, str):
raise ValueError("Param rename key must be a string if set")
rename = rename.strip()
if not rename or not all(c == "_" or c.isalnum() for c in rename):
raise ValueError(
"Param rename keys must consist of one or more alphanumeric or underscore characters"
)
return f"{key}_{rename}"
def go(self, *, key, value, params_opt=None, parameters=""):
MODULES.initialize()
if not self.OCS_PARAM_INPUT_TYPES[key](value):
if not self.OCS_PARAM_TYPES[key](value):
raise ValueError(f"CSamplerParam: Bad value type for key {key}")
if parameters:
extra_params = yaml.safe_load(parameters)
if extra_params is not None:
if not isinstance(extra_params, dict):
raise ValueError("Parameters must be a JSON or YAML object")
key = self.get_renamed_key(key, extra_params)
else:
extra_params = None
params = ParamGroup(items={}) if params_opt is None else params_opt.clone()
@@ -431,7 +349,7 @@ class ParamNode(metaclass=IntegratedNode):
return (params,)
class MultiParamNode(ParamNode, metaclass=IntegratedNode):
class MultiParamNode:
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."
@@ -446,7 +364,7 @@ class MultiParamNode(ParamNode, metaclass=IntegratedNode):
@classmethod
def INPUT_TYPES(cls):
param_keys = (
("", *ParamNode.OCS_PARAM_INPUT_TYPES.keys()),
("", *ParamNode.OCS_PARAM_TYPES.keys()),
{
"tooltip": "Used to set the type of custom parameter.",
},
@@ -460,30 +378,26 @@ class MultiParamNode(ParamNode, metaclass=IntegratedNode):
"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": "",
"placeholder": """\
"default": """\
# 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}": (
WILDCARD_PARAM,
ParamNode.WC,
{
"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,
"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.",
},
)
for idx in range(1, cls.PARAM_COUNT + 1)
@@ -491,7 +405,6 @@ class MultiParamNode(ParamNode, metaclass=IntegratedNode):
}
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)
@@ -506,18 +419,17 @@ class MultiParamNode(ParamNode, metaclass=IntegratedNode):
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_INPUT_TYPES[key](value):
if not ParamNode.OCS_PARAM_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)
params[key] = value
extra = extra_params.get(str(idx))
if extra is not None:
params[f"{key}.params"] = extra
return (params,)
class SimpleRestartSchedule(metaclass=IntegratedNode):
class SimpleRestartSchedule:
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."
@@ -550,9 +462,8 @@ class SimpleRestartSchedule(metaclass=IntegratedNode):
"schedule": (
"STRING",
{
"default": "",
"placeholder": """\
# YAML or JSON restart schedule. Example:
"default": """\
# YAML or JSON restart schedule
# Every 5 steps, jump back 3 steps
- [5, -3]
# Jump to schedule item 0
@@ -566,8 +477,7 @@ class SimpleRestartSchedule(metaclass=IntegratedNode):
},
}
def go(self, *, sigmas, start_step=0, schedule=""):
MODULES.initialize()
def go(self, *, sigmas, start_step=0, schedule="[]"):
if schedule:
parsed_schedule = yaml.safe_load(schedule)
if parsed_schedule is not None:
@@ -580,12 +490,12 @@ class SimpleRestartSchedule(metaclass=IntegratedNode):
return (Restart.simple_schedule(sigmas, start_step, parsed_schedule),)
class ModelSetMaxSigmaNode(metaclass=IntegratedNode):
class ModelSetMaxSigmaNode:
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, the main use case is only for generating sigmas.",
"Patched model. Can be connected to a scheduler node (i.e. BasicScheduler). Generally NOT recommended to connect to an actual sampler.",
)
FUNCTION = "go"
@@ -603,7 +513,7 @@ class ModelSetMaxSigmaNode(metaclass=IntegratedNode):
"mode": (
("recalculate", "simple_multiply"),
{
"tooltip": "Mode to use when setting sigmas in the patched model. Recalculate should generally be more accurate.",
"tooltip": "Mode use for setting sigmas in the patched model. Recalculate should generally be more accurate.",
},
),
"sigma_max": (
@@ -614,7 +524,7 @@ class ModelSetMaxSigmaNode(metaclass=IntegratedNode):
"max": 10000.0,
"step": 0.01,
"round": False,
"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).",
"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).",
},
),
"fake_sigma_min": (
@@ -625,14 +535,13 @@ class ModelSetMaxSigmaNode(metaclass=IntegratedNode):
"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 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).",
"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).",
},
),
}
}
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"):
@@ -673,441 +582,12 @@ class ModelSetMaxSigmaNode(metaclass=IntegratedNode):
"ModelSetMaxSigma: Invalid fake_min_sigma value, result max <= min"
)
model.add_object_patch("model_sampling", ms)
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}"
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}"
)
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,)
class ExpressionFilteredLatentOperation:
EXTENDED_LATENT_OPERATION = True
def __init__(
self,
*,
ocs_filter: Filter,
latent_refs: dict[str, torch.Tensor] | None = None,
) -> None:
self.filter = ocs_filter
self.latent_refs = latent_refs if latent_refs is not None else {}
def __call__(
self,
latent: torch.Tensor,
*,
sigma: float | torch.Tensor | None = None,
**kwargs: dict,
) -> torch.Tensor:
refs = FilterRefs(
kvs=kwargs
| {
"sigma": sigma.clone() if isinstance(sigma, torch.Tensor) else sigma,
"sigma_float": sigma.max().item()
if isinstance(sigma, torch.Tensor)
else sigma,
}
| {k: v.to(latent, copy=True) for k, v in self.latent_refs.items()}
)
return self.filter.apply(latent, refs=refs)
class ExpressionFilteredLatentOperationNode:
DESCRIPTION = "TBD"
FUNCTION = "go"
RETURN_TYPES = ("LATENT_OPERATION",)
@classmethod
def INPUT_TYPES(cls):
MODULES.initialize()
return {
"required": {
"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.",
},
),
},
"optional": {
"latent_ref_1_opt": ("LATENT",),
"latent_ref_2_opt": ("LATENT",),
"latent_ref_3_opt": ("LATENT",),
},
}
def go(
self,
*,
yaml_config: str,
latent_ref_1_opt: dict | None = None,
latent_ref_2_opt: dict | None = None,
latent_ref_3_opt: dict | None = None,
) -> tuple:
config = yaml.safe_load(yaml_config)
if isinstance(config, str):
config = {"filter": {"final": config}}
elif 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 TypeError("Bad type for filter definition, must be object")
latent_refs = {
k: v["samples"].to(device="cpu", dtype=torch.float32, copy=True)
for k, v in (
("latent_ref_1", latent_ref_1_opt),
("latent_ref_2", latent_ref_2_opt),
("latent_ref_3", latent_ref_3_opt),
)
if v is not None
}
ocs_filter = make_filter(filter_def)
return (
ExpressionFilteredLatentOperation(
ocs_filter=ocs_filter, latent_refs=latent_refs
),
)
class ExpressionFilteredModelPatchNode:
DESCRIPTION = "TBD"
FUNCTION = "go"
RETURN_TYPES = ("MODEL",)
@classmethod
def INPUT_TYPES(cls):
MODULES.initialize()
return {
"required": {
"model": ("MODEL",),
"patch_mode": (
("apply_model", "pre_cfg", "post_cfg", "cfg", "denoise_mask"),
{"default": "apply_model"},
),
"existing_patch_mode": (
("normal", "extract", "extract_split", "extract_sequence"),
{
"default": "normal",
"tooltip": "Modes:\n"
"normal: Replaces apply_model or cfg patches, appends for pre_cfg and post_cfg.\n"
"extract: Removes the existing patches and passes old_result with the output from existing patches.\n"
"extract_split: Same as extract except you'll get tuple of results for each existing patch (apply_model and cfg will always be length 1).\n"
"extract_sequence: Like extract_split except existing patches do not see each other's effects.",
},
),
"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.",
},
),
},
"optional": {
"latent_ref_1_opt": ("LATENT",),
"latent_ref_2_opt": ("LATENT",),
"latent_ref_3_opt": ("LATENT",),
},
}
def go(
self,
*,
yaml_config: str,
model: object,
patch_mode: str,
existing_patch_mode: str,
latent_ref_1_opt: dict | None = None,
latent_ref_2_opt: dict | None = None,
latent_ref_3_opt: dict | None = None,
) -> tuple:
config = yaml.safe_load(yaml_config)
if isinstance(config, str):
config = {"filter": {"final": config}}
elif 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 TypeError("Bad type for filter definition, must be object")
latent_refs = {
k: v["samples"].to(device="cpu", dtype=torch.float32, copy=True)
for k, v in (
("latent_ref_1", latent_ref_1_opt),
("latent_ref_2", latent_ref_2_opt),
("latent_ref_3", latent_ref_3_opt),
)
if v is not None
}
ocs_filter = make_filter(filter_def)
model = model.clone()
mode_keys = {
"post_cfg": "sampler_post_cfg_function",
"pre_cfg": "sampler_pre_cfg_function",
"apply_model": "model_function_wrapper",
"cfg": "sampler_cfg_function",
"denoise_mask": "denoise_mask_function",
}
key = mode_keys.get(patch_mode)
if key is None:
raise ValueError(f"Bad mode: {patch_mode}")
if existing_patch_mode != "normal":
old_handlers = model.model_options.pop(key, None)
if old_handlers is None:
old_handlers = ()
else:
old_handlers = ()
def get_refs(*args, **kwargs) -> FilterRefs:
old_results = []
if patch_mode in {"pre_cfg", "post_cfg", "cfg"}:
argdict = args[0]
elif patch_mode == "denoise_mask":
argdict = {
"sigma": args[0],
"denoise_mask": args[1].clone(),
"sigmas": kwargs["extra_options"]["sigmas"].clone(),
}
elif patch_mode == "apply_model":
argdict = args[1] | {"apply_function": args[0]}
else:
raise ValueError(f"Bad patch mode: {patch_mode}")
if old_handlers:
ridx = 0 if existing_patch_mode == "extract_sequence" else -1
if patch_mode in {"cfg", "denoise_mask", "apply_model"}:
old_results = (old_handlers[0](*args, **kwargs),)
elif patch_mode == "pre_cfg":
old_results = [argdict["conds_out"]]
for hf in old_handlers:
result = hf(argdict | {"conds_out": old_results[ridx]}).copy()
if len(old_results) > 1 and existing_patch_mode != "extract":
old_results[1] = result
else:
old_results.append(result)
old_results = old_results[1:]
elif patch_mode == "post_cfg":
old_results = [argdict["denoised"].clone()]
for hf in old_handlers:
result = hf(argdict | {"denoised": old_results[ridx].clone()})
if len(old_results) > 1 and existing_patch_mode != "extract":
old_results[1] = result
else:
old_results.append(result)
old_results = old_results[1:]
kvs = {
"sigma": argdict["sigma"].clone(),
"sigma_float": argdict["sigma"].max().item(),
"old_results": tuple(old_results),
}
if patch_mode in {"pre_cfg", "post_cfg", "cfg"}:
kvs |= {
"x": argdict["input"].clone(),
"cfg_scale": argdict["cond_scale"],
}
if patch_mode in {"post_cfg", "cfg"}:
kvs["cond"] = argdict["cond_denoised"].clone()
uncond = argdict.get("uncond_denoised", None)
kvs["uncond"] = uncond if uncond is None else uncond.clone()
if patch_mode == "post_cfg":
kvs["denoised"] = argdict["denoised"].clone()
else:
conds_out = argdict["conds_out"]
kvs["cond"] = conds_out[0].clone()
kvs["uncond"] = (
conds_out[1].clone()
if len(conds_out) > 1 and conds_out[1] is not None
else None
)
kvs["conds_out"] = list(conds_out)
elif patch_mode == "denoise_mask":
kvs |= {
"sigmas": argdict["sigmas"],
"denoise_mask": argdict["denoise_mask"],
}
elif patch_mode == "apply_model":
kvs |= {
"x": argdict["input"].clone(),
"cond_or_uncond": argdict["cond_or_uncond"].clone(),
}
else:
raise ValueError(f"Bad patch mode: {patch_mode}")
if patch_mode in {"pre_cfg", "post_cfg", "cfg", "apply_model"}:
latent_in = kvs["x"]
else:
latent_in = kvs["denoise_mask"]
kvs |= {k: v.to(latent_in) for k, v in latent_refs.items()}
return FilterRefs(kvs=kvs)
def model_patch(*args, **kwargs):
refs = get_refs(*args, **kwargs)
if patch_mode == "apply_model":
def fallback_apply_model():
old_results = refs.kvs["old_results"]
if old_results:
return old_results[-1]
return args[0](
args[1]["input"], args[1]["timestep"], **args[1]["c"]
)
else:
fallback_apply_model = None
if not ocs_filter.check_applies(refs):
old_results = refs.kvs["old_results"]
if old_results:
return old_results[-1]
if patch_mode == "pre_cfg":
return args[0]["conds_out"]
if patch_mode == "post_cfg":
return args[0]["denoised"]
if patch_mode == "cfg":
return args[0]["cond"]
if patch_mode == "denoise_mask":
return args[1]
if patch_mode == "apply_model":
return fallback_apply_model()
raise ValueError(f"Bad patch mode: {patch_mode}")
if patch_mode in {"pre_cfg", "post_cfg", "cfg", "apply_model"}:
latent_in = refs.kvs["x"]
else:
latent_in = refs.kvs["denoise_mask"]
result = ocs_filter.apply(latent_in, refs=refs)
if patch_mode == "apply_model" and result is None:
return fallback_apply_model()
if patch_mode == "pre_cfg":
return list(result)
return result
if patch_mode == "pre_cfg":
model.set_model_sampler_pre_cfg_function(model_patch)
elif patch_mode == "post_cfg":
model.set_model_sampler_post_cfg_function(model_patch)
elif patch_mode == "cfg":
model.set_model_sampler_cfg_function(model_patch)
elif patch_mode == "denoise_mask":
model.set_model_denoise_mask_function(model_patch)
elif patch_mode == "apply_model":
model.set_model_unet_function_wrapper(model_patch)
else:
raise ValueError(f"Bad patch mode: {patch_mode}")
return (model,)
__all__ = (
"SamplerNode",
"GroupNode",
+44 -187
View File
@@ -1,68 +1,11 @@
import gc
import math
import random
import numpy as np
import scipy
import torch
from tqdm import tqdm
from .filtering import Filter, make_filter
from .utils import fallback, scale_noise
try:
from .triton_lsa import (
assignments_to_indices,
batch_linear_assignment,
batch_linear_assignment_shuffled,
)
HAVE_TRITON = True
except Exception:
HAVE_TRITON = False
def linear_sum_assignment(
cost: torch.Tensor,
*,
maximize: bool = False,
use_triton: bool = False,
split_batch: int = 0,
**kwargs: dict,
) -> tuple[np.ndarray, np.ndarray] | tuple[torch.Tensor, torch.Tensor]:
if not use_triton or not HAVE_TRITON or not cost.is_cuda:
cost = cost.half().cpu()
return scipy.optimize.linear_sum_assignment(cost, maximize=maximize)
ndim = cost.ndim
orig_shape = cost.shape
if ndim == 2:
do_split = split_batch > 1 and all(
(sz / split_batch).is_integer() for sz in orig_shape
)
if do_split:
cost = cost.reshape(
split_batch, orig_shape[0] // split_batch, orig_shape[1] // split_batch
)
else:
cost = cost.unsqueeze(0)
tqdm.write(
f"TRITON LAP: maximize={maximize}, orig cost shape={orig_shape}, cost shape={cost.shape}, cost dtype={cost.dtype}",
)
if not cost.is_contiguous():
cost = cost.contiguous()
fun = (
batch_linear_assignment
if "generator" not in kwargs
else batch_linear_assignment_shuffled
)
assignments = fun(cost, maximize=maximize, **kwargs)
row_ind, col_ind = assignments_to_indices(assignments)
if ndim == 2:
row_ind = row_ind.reshape(-1, row_ind.shape[-1])
col_ind = col_ind.reshape(-1, col_ind.shape[-1])
# row_ind, col_ind = row_ind.squeeze(0), col_ind.squeeze(0)
tqdm.write(f"Ran LAP kernel: {assignments.shape}, {row_ind.shape}, {col_ind.shape}")
return row_ind, col_ind
from .utils import scale_noise, fallback
class ImmiscibleNoise(Filter):
@@ -72,22 +15,15 @@ class ImmiscibleNoise(Filter):
"size": 0,
"batching": "channel",
"maximize": False,
"distance_scale": 0.0,
"distance_scale_ref": None,
"abs_mode": False,
"abs_distance_mode": False,
"use_triton": False,
# Only honored in Triton mode.
"split_batch": 0,
"generator": None,
}
def __call__(self, noise_sampler, x_ref, *, refs=None):
if self.size == 0 or self.strength == 0 or not self.check_applies(refs):
if 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(size))),
torch.cat(tuple(noise_sampler() for _ in range(self.size)))
if self.size > 0
else noise_sampler(),
default_ref=x_ref,
refs=refs,
output_shape=x_ref.shape,
@@ -96,107 +32,52 @@ 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[offset:]), self.batch(ref_latent)),
output_shape,
self.immiscible(self.batch(latent), 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.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:])
return latent.view(sz[0] * sz[1], *sz[2:])
if self.batching == "row":
return latent.reshape(math.prod(sz[:-1]), sz[-1])
return latent.view(sz[0] * sz[1] * sz[2], sz[3])
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.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)
return latent.view(*sz[:2], sz[3], sz[2]).permute(0, 1, 3, 2)
return latent.view(*sz)
# 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,
):
# 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):
# "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303
# Minimize latent-noise pairs over a batch
batch = latent.shape[0]
out_latent = fallback(out_latent, latent)
ref_latent = ref_latent.detach().clone()
if self.abs_mode:
ref_latent = ref_latent.abs()
latent = latent.abs()
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
if self.abs_distance_mode:
dist = dist.abs_()
del ref_latent_expanded, latent_expanded
dist = dist.mean(tuple(range(2, dist.ndim)))
else:
distance_scale_ref = fallback(self.distance_scale_ref, self.distance_scale)
dist = distance_scale_ref * ref_latent.flatten(start_dim=1).unsqueeze(
1
) - self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0)
if self.abs_distance_mode:
dist = dist.abs_()
dist = torch.linalg.vector_norm(dist, dim=2)
try:
assign_mat = linear_sum_assignment(
dist,
maximize=self.maximize,
use_triton=self.use_triton,
split_batch=self.split_batch,
generator=self.generator,
)
except ValueError as exc:
tqdm.write(f"OCS: Immiscible: Failed due to exception: {exc}")
return None if return_idxs else out_latent[: ref_latent.shape[0]]
return assign_mat if return_idxs else out_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,
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()
try:
assign_mat = scipy.optimize.linear_sum_assignment(
dist, 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]]
class NoiseSamplerCache:
@@ -209,11 +90,10 @@ class NoiseSamplerCache:
*,
normalize_noise=True,
cpu_noise=True,
batch_size=1,
caching=True,
batch_size=32,
caching=False,
cache_reset_interval=9999,
set_seed=True,
seed_offset=1,
set_seed=False,
scale=1.0,
normalize_dims=(-3, -2, -1),
immiscible=None,
@@ -223,7 +103,7 @@ class NoiseSamplerCache:
self.x = x
self.mega_x = None
self.seed = seed
self.seed_offset = seed_offset
self.seed_offset = 0
self.min_sigma = min_sigma
self.max_sigma = max_sigma
self.cache = {}
@@ -243,11 +123,6 @@ 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 = {}
@@ -285,14 +160,10 @@ class NoiseSamplerCache:
size,
sigma,
sigma_next,
*,
immiscible=None,
sigmas=None,
):
size = min(size, self.batch_size)
if immiscible is None:
immiscible = self.immiscible
cache_key = (nsobj, size, hash(immiscible))
cache_key = (nsobj, size)
if self.caching:
noise_sampler = self.cache.get(cache_key)
if noise_sampler:
@@ -302,26 +173,14 @@ class NoiseSamplerCache:
curr_x = self.mega_x[: self.x.shape[0] * size, ...]
if nsobj is None:
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
def ns(_s, _sn, *_unused, **_unusedkwargs):
return torch.randn_like(curr_x)
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,
sigma_min,
sigma_max,
self.min_sigma,
self.max_sigma,
seed=curr_seed,
normalized=False,
cpu=self.cpu_noise,
@@ -330,10 +189,10 @@ 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,
curr_sigma_next,
*_unused,
out_hw=(orig_h, orig_w),
**_unusedkwargs,
@@ -344,9 +203,7 @@ class NoiseSamplerCache:
f"Noise size mismatch: {out_hw} vs {(orig_h, orig_w)}"
)
if remain < 1:
curr_sigma = fallback(curr_sigma, sigma)
curr_sigma_next = fallback(curr_sigma_next, sigma_next)
noise = self.scale_noise(ns(curr_sigma, curr_sigma_next)).view(
noise = self.scale_noise(ns(sigma, sigma_next)).view(
size,
*self.x.shape,
)
+10 -112
View File
@@ -1,57 +1,8 @@
from __future__ import annotations
from typing import NamedTuple
import torch
# from tqdm import tqdm
class RestartScaleFactors(NamedTuple):
latent_scale: float
noise_scale: float
@classmethod
def build(
cls,
sigma_from: float | torch.Tensor,
sigma_to: float | torch.Tensor,
*,
is_flow: bool,
) -> RestartScaleFactors:
if isinstance(sigma_from, torch.Tensor):
sigma_from = sigma_from.max().item()
if isinstance(sigma_to, torch.Tensor):
sigma_to = sigma_to.max().item()
if not is_flow:
return cls(
1.0,
max(0.0, (sigma_to**2 - sigma_from**2)) ** 0.5,
)
alpha_from = 1.0 - sigma_from
alpha_to = 1.0 - sigma_to
if alpha_to <= 0:
latent_scale = 0.0
noise_scale = sigma_to
else:
latent_scale = alpha_to / alpha_from
noise_scale = (
max(0.0, (sigma_to**2) - (latent_scale * sigma_from) ** 2) ** 0.5
)
return cls(latent_scale, noise_scale)
class Restart:
def __init__(
self,
*,
s_noise=1.0,
custom_noise=None,
immiscible=False,
normalized=True,
normalize_dims: tuple[int, ...] | None = None,
is_flow=False,
):
def __init__(self, *, s_noise=1.0, custom_noise=None, immiscible=False):
from .noise import ImmiscibleNoise
self.s_noise = s_noise
@@ -59,9 +10,6 @@ class Restart:
immiscible = ImmiscibleNoise(**immiscible)
self.immiscible = immiscible
self.custom_noise = custom_noise
self.normalized = normalized
self.normalize_dims = normalize_dims
self.is_flow = is_flow
def get_noise_sampler(self, nsc):
return nsc.make_caching_noise_sampler(
@@ -82,64 +30,23 @@ class Restart:
last_sigma = sigma
return sigmas
def split_sigmas(self, sigmas: torch.Tensor):
def split_sigmas(self, sigmas):
prev_seg = None
while len(sigmas) > 1:
seg = self.get_segment(sigmas)
sigmas = sigmas[len(seg) :]
if prev_seg is not None and seg[0] > prev_seg[-1]:
scale_factors = RestartScaleFactors.build(
sigma_from=prev_seg[-1], sigma_to=seg[0], is_flow=self.is_flow
)
noise_scale = self.get_noise_scale(prev_seg[-1], seg[0])
else:
scale_factors = None
noise_scale = 0.0
prev_seg = seg
yield (scale_factors, seg)
yield (noise_scale, seg)
def get_noise_scale(
self, s_min: float | torch.Tensor, s_max: float | torch.Tensor
) -> float:
def get_noise_scale(self, s_min, s_max):
result = (s_max**2 - s_min**2) ** 0.5
if isinstance(result, torch.Tensor):
return result.item()
return result
def add_noise(
self,
x: torch.Tensor,
sigma_from: float,
sigma_to: float,
*,
nsc,
refs,
scale_factors: RestartScaleFactors | None = None,
in_place: bool = False,
) -> torch.Tensor:
if self.is_flow:
sigma_from = min(1.0, max(0.0, sigma_from))
sigma_to = min(1.0, max(0.0, sigma_to))
if sigma_from >= sigma_to:
raise ValueError(
f"sigma_from ({sigma_from:.4f}) must be less than sigma_to ({sigma_to:.4f})"
)
scale_factors = scale_factors or RestartScaleFactors.build(
sigma_from, sigma_to, is_flow=self.is_flow
)
ns = self.get_noise_sampler(nsc)
sigma_empty = nsc.min_sigma * 0
noise = nsc.scale_noise(
ns(sigma_empty + sigma_from, sigma_empty + sigma_to, refs=refs),
normalized=self.normalized,
normalize_dims=self.normalize_dims,
)
noise *= scale_factors.noise_scale * self.s_noise
if scale_factors.latent_scale != 1.0:
x = (
x.mul_(scale_factors.latent_scale)
if in_place
else scale_factors.latent_scale * x
)
return noise.add_(x)
result = result.item()
return result * self.s_noise
def __repr__(self):
return f"<Restart: s_noise={self.s_noise:.04}, immiscible={self.immiscible}>"
@@ -170,27 +77,18 @@ class Restart:
raise ValueError("Schedule jump index out of range")
sched_idx = item
continue
sched_frac = round(sig_idx - int(sig_idx), ndigits=5)
sig_idx = int(sig_idx if sched_frac == 0 else sig_idx + 1)
if sig_idx >= siglen or sig_idx < 0:
break
interval, jump = item
chunk = siglist[sig_idx : sig_idx + interval + 1]
if sched_frac != 0:
chunk[0] -= (chunk[0] - chunk[1]) * (1.0 - sched_frac)
# print(f"{out} + {chunk}")
out += chunk
sig_idx += interval + jump
if jump >= 0:
sig_idx += 1
sig_idx += interval + jump
sched_idx += 1
sched_frac = round(sig_idx - int(sig_idx), ndigits=5)
sig_idx = int(sig_idx if sched_frac == 0 else sig_idx + 1)
if sig_idx < siglen and sig_idx >= 0:
chunk = siglist[sig_idx:]
if sched_frac != 0:
chunk[0] -= (chunk[0] - chunk[1]) * (1.0 - sched_frac)
out += chunk
out += siglist[sig_idx:]
if out[-1] > siglist[-1]:
out.append(siglist[-1])
return torch.tensor(out).to(sigmas)
+29 -45
View File
@@ -1,12 +1,13 @@
import torch
from tqdm.auto import trange
from .filtering import FILTER_HANDLERS, FilterRefs
from .model import OCSModel
from .filtering import FILTER_HANDLERS
from .model import ModelCallCache
from .noise import NoiseSamplerCache
from .restart import Restart
from .substep_merging import MERGE_SUBSTEPS_CLASSES
from .substep_sampling import SamplerState
from .substep_merging import MERGE_SUBSTEPS_CLASSES
from .restart import Restart
def find_merge_sampler(merge_samplers, ss) -> object | None:
@@ -43,13 +44,14 @@ 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}")
restart = Restart(
s_noise=restart_params.get("s_noise", 1.0),
custom_noise=copts.get("restart_custom_noise"),
immiscible=restart_params.get("immiscible", False),
)
ss = SamplerState(
OCSModel(
ModelCallCache(
model,
x,
x.new_ones((x.shape[0],)),
@@ -61,21 +63,11 @@ def composable_sampler(
extra_args,
noise_sampler=noise_sampler,
callback=callback,
eta=eta if eta != 1.0 else copts.get("eta", 1.0),
s_noise=s_noise if s_noise != 1.0 else copts.get("s_noise", 1.0),
eta=eta if eta != 1.0 else copts["eta"],
s_noise=s_noise if s_noise != 1.0 else copts["s_noise"],
reta=copts.get("reta", 1.0),
disable_status=disable,
)
restart = Restart(
s_noise=restart_params.get("s_noise", 1.0),
custom_noise=restart_custom_noise,
immiscible=restart_params.get("immiscible", False),
normalized=restart_params.get("normalized", True),
normalize_dims=restart_params.get("normalize_dims"),
is_flow=ss.model.is_rectified_flow,
)
groups = copts["_groups"]
merge_samplers = tuple(
MERGE_SUBSTEPS_CLASSES[g.merge_method](ss, g) for g in groups.items
@@ -83,23 +75,18 @@ def composable_sampler(
nsc = NoiseSamplerCache(
x,
extra_args.get("seed", 42),
sigmas[sigmas > 0].min(),
sigmas.max(),
sigmas[-1],
sigmas[0],
**copts.get("noise", {}),
)
ss.noise = nsc
sigma_chunks = (
tuple(restart.split_sigmas(sigmas)) if restart_enabled else ((None, sigmas),)
)
sigma_chunks = tuple(restart.split_sigmas(sigmas))
step_count = sum(len(chunk) - 1 for _noise, chunk in sigma_chunks)
ss.total_steps = step_count
step = 0
restart_snoise = copts.get("restart_s_noise", 1.0)
with trange(step_count, disable=ss.disable_status) as pbar:
for chunk_idx, (scale_factors, chunk_sigmas) in enumerate(sigma_chunks):
if step != 0 and scale_factors is not None:
prev_refs = FilterRefs(
{f"pre_restart_{k}": v for k, v in ss.refs.items()}
)
for noise_scale, chunk_sigmas in sigma_chunks:
ss.sigmas = chunk_sigmas
ss.update(0, step=step, substep=0)
if step != 0:
@@ -108,25 +95,22 @@ def composable_sampler(
ss.hist.reset()
for ms in merge_samplers:
ms.reset()
nsc.min_sigma, nsc.max_sigma = (
chunk_sigmas[-1].clone(),
chunk_sigmas[0].clone(),
)
if step != 0 and scale_factors is not None:
x = restart.add_noise(
x,
sigma_from=sigma_chunks[chunk_idx - 1][1][-1].item(),
sigma_to=chunk_sigmas[0].item(),
scale_factors=scale_factors,
nsc=nsc,
refs=prev_refs | ss.refs,
in_place=True,
nsc.min_sigma, nsc.max_sigma = chunk_sigmas[-1], chunk_sigmas[0]
if step != 0 and noise_scale != 0:
restart_ns = restart.get_noise_sampler(nsc)
x += nsc.scale_noise(
restart_ns(refs=ss.refs),
noise_scale * restart_snoise,
)
del prev_refs
del restart_ns
for idx in range(len(chunk_sigmas) - 1):
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:
+2214
View File
File diff suppressed because it is too large Load Diff
-20
View File
@@ -1,20 +0,0 @@
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")
-536
View File
@@ -1,536 +0,0 @@
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"<SS({self.name}): s_noise={self.s_noise}, eta={self.eta}>"
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
-548
View File
@@ -1,548 +0,0 @@
import inspect
import math
import typing
import comfy
import torch
from .. import expression as expr
from .. import filtering
from ..utils import fallback
from .base import (
SingleStepSampler,
StepSamplerContext,
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:
bi_elements = math.prod(x.shape[1:])
denoised_new = denoised_new.reshape(x.shape[0], -1)[
:, :bi_elements
].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,
)
-684
View File
@@ -1,684 +0,0 @@
import comfy
import torch
from comfy.k_diffusion.sampling import get_ancestral_step
from tqdm import tqdm
from .. import filtering
from .base import (
DPMPPStepMixin,
HistorySingleStepSampler,
ReversibleSingleStepSampler,
SingleStepSampler,
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)
# SEEDS-2 - Stochastic Explicit Exponential Derivative-free Solvers (VP Data Prediction) stage 2.
# arXiv: https://arxiv.org/abs/2305.14267 (NeurIPS 2023)
# Implementation referenced from ComfyUI.
class Seeds2Step(SingleStepSampler, DPMPPStepMixin):
name = "seeds_2"
self_noise = 3
model_calls = 1
allow_alt_cfgpp = False
uses_alt_noise = True
def __init__(self, *args, r=0.5, **kwargs):
super().__init__(*args, **kwargs)
self.r = r
s2_options = self.options.get("seeds_2", {})
sigma_blend_mode = s2_options.get("sigma_blend_mode", "lerp").strip()
self.sigma_blend_function = (
filtering.BLENDING_MODES[sigma_blend_mode]
if sigma_blend_mode != "lerp"
else torch.lerp
)
denoised_blend_mode = s2_options.get("denoised_blend_mode", "lerp").strip()
self.denoised_blend_function = (
filtering.BLENDING_MODES[denoised_blend_mode]
if denoised_blend_mode != "lerp"
else torch.lerp
)
self.disable_stage2_eta = bool(s2_options.get("disable_stage2_eta", False))
stage2_stage1_noise_blend_mode = s2_options.get(
"stage2_stage1_noise_blend_mode", "lerp"
).strip()
self.stage2_stage1_noise_blend_function = (
filtering.BLENDING_MODES[stage2_stage1_noise_blend_mode]
if stage2_stage1_noise_blend_mode != "lerp"
else torch.lerp
)
self.stage2_stage1_noise_ratio = s2_options.get(
"stage2_stage1_noise_ratio", 1.0
)
self.stage1_s_noise = s2_options.get("stage1_s_noise", 1.0)
self.stage2_s_noise = s2_options.get("stage2_s_noise", 1.0)
self.stage2_sigma_scale = s2_options.get("stage2_sigma_scale", 1.0)
def step(self, x: torch.Tensor):
ss = self.ss
sigma = ss.sigma.to(dtype=torch.float64)
sigma_next = ss.sigma_next.to(dtype=torch.float64)
denoised = ss.denoised
t_one = ss.sigma * 0 + 1.0
r, eta = self.r, self.get_dyn_eta()
fac = 1 / (2 * r)
lambda_s = ss.sigma_to_half_log_snr(sigma=sigma)
lambda_t = ss.sigma_to_half_log_snr(sigma=sigma_next)
h = lambda_t - lambda_s
h_eta = h * (eta + 1.0)
lambda_s_1 = self.sigma_blend_function(
lambda_s.unsqueeze(0), lambda_t.unsqueeze(0), r
).squeeze(0)
sigma_s_1 = ss.half_log_snr_to_sigma(lambda_s_1)
alpha_s_1 = sigma_s_1 * lambda_s_1.exp()
alpha_t = sigma_next * lambda_t.exp()
s1_x_mult = sigma_s_1 / sigma * (-r * h * eta).exp()
s1_denoised_mult = alpha_s_1 * (-r * h_eta).expm1()
x_2 = (
s1_x_mult.to(dtype=x.dtype) * x
- s1_denoised_mult.to(dtype=x.dtype) * denoised
)
if eta != 0:
s1_noise_mult = (-2 * r * h * eta).expm1().neg().sqrt()
sde_noise1 = yield from self.result(
x_2 * 0,
s1_noise_mult.to(dtype=x.dtype),
sigma=ss.sigma,
sigma_next=sigma_s_1.to(dtype=x.dtype),
noise_sampler=self.alt_noise_sampler,
final=False,
)
x_2 += sde_noise1 * (sigma_s_1.to(dtype=x.dtype) * self.stage1_s_noise)
denoised_2 = self.call_model(
x_2, (sigma_s_1 * self.stage2_sigma_scale).to(dtype=x.dtype), call_index=1
).denoised
denoised_d = self.denoised_blend_function(denoised, denoised_2, fac)
if self.disable_stage2_eta:
eta = 0.0
h_eta = h
s2_x_mult = sigma_next / sigma * (-h * eta).exp()
s2_denoised_mult = alpha_t * h_eta.neg().expm1()
x_curr = s2_x_mult.to(dtype=x.dtype) * x
x_curr -= s2_denoised_mult.to(dtype=x.dtype) * denoised_d
if eta == 0:
return (yield from self.result(x_curr))
s2_s1_nr = self.stage2_stage1_noise_ratio
segment_factor = ((r - 1.0) * h * eta).to(dtype=x.dtype)
s2_noise_mult = (segment_factor * 2.0).expm1().neg() ** 0.5
sde_noise2_raw = yield from self.result(
x_curr * 0,
t_one,
sigma=sigma_s_1.to(dtype=x.dtype),
sigma_next=ss.sigma_next,
final=False,
)
sde_noise2 = sde_noise2_raw * s2_noise_mult.to(dtype=x.dtype)
if s2_s1_nr != 1.0:
# print(
# f"\n\nBLENDING: {s2_s1_nr:.4f}, {s1_noise_mult.item():.4f}, {s2_noise_mult.item():.4f}"
# )
sde_noise1 = self.stage2_stage1_noise_blend_function(
(
yield from self.result(
x_curr * 0,
t_one,
sigma=sigma_s_1.to(dtype=x.dtype),
sigma_next=ss.sigma_next,
final=False,
)
)
* s1_noise_mult.to(dtype=x.dtype),
sde_noise1,
s2_s1_nr,
)
sde_noise1 *= segment_factor.exp()
sde_noise2 += sde_noise1
sde_noise2 *= ss.sigma_next * self.stage2_s_noise
x_curr += sde_noise2
yield from self.result(
x_curr, noise_scale=ss.sigma * 0, sigma_down=ss.sigma_next
)
registry.add(
DEISStep,
DPMPP2MSDEStep,
DPMPP2MStep,
DPMPP3MSDEStep,
DPMPPSDEStep,
EulerStep,
HeunPP2Step,
IPNDMStep,
IPNDMVStep,
GradientEstimationStep,
DPM2Step,
DPMPP2SStep,
RESMultistepStep,
Seeds2Step,
)
-755
View File
@@ -1,755 +0,0 @@
# 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,
)
-131
View File
@@ -1,131 +0,0 @@
# 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)
-29
View File
@@ -1,29 +0,0 @@
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", 9999)
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,
)
-39
View File
@@ -1,39 +0,0 @@
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}
-49
View File
@@ -1,49 +0,0 @@
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)
-304
View File
@@ -1,304 +0,0 @@
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)
-118
View File
@@ -1,118 +0,0 @@
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)
-132
View File
@@ -1,132 +0,0 @@
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)
-192
View File
@@ -1,192 +0,0 @@
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)
+124 -362
View File
@@ -5,12 +5,10 @@ import tqdm
from . import expression as expr
from . import utils
from .filtering import FILTER_HANDLERS, FilterRefs, make_filter
from .noise import ImmiscibleNoise
from .filtering import make_filter, FilterRefs
from .restart import Restart
from .step_samplers import STEP_SAMPLERS
from .step_samplers.base import StepSamplerContext
from .substep_sampling import StepSamplerChain
from .utils import check_time, fallback
@@ -22,7 +20,6 @@ class MergeSubstepsSampler:
STEP_SAMPLERS[sitem["step_method"]](**sitem) for sitem in group.items
)
options = group.options.copy()
self.group = group
self.time_mode = group.time_mode
self.time_start = group.time_start
self.time_end = group.time_end
@@ -36,10 +33,6 @@ class MergeSubstepsSampler:
self.pre_filter = None if pre_filter is None else make_filter(pre_filter)
self.post_filter = None if post_filter is None else make_filter(post_filter)
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):
@@ -60,51 +53,35 @@ 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
x = self.pre_filter.apply(x, refs=fallback(ss, self.ss).refs)
ss.noise.update_x(x)
return x
ss = fallback(ss, self.ss)
return self.pre_filter.apply(x, refs=fallback(ss, self.ss).refs)
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})
x = self.post_filter.apply(x, refs=refs)
ss.noise.update_x(x)
return x
return self.post_filter.apply(x, refs=refs)
def __call__(self, x):
orig_x = x
x = self.step_input(x)
if self.afs_start_step <= self.ss.step <= self.afs_end_step:
x = self.afs_step(x)
else:
x = self.step(x)
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
def substep(self, x, sampler):
sg = sampler(x)
def substep(self, x, sampler, ss=None):
sg = sampler(x, fallback(ss, self.ss))
yield from utils.step_generator(sg, get_next=lambda sr: sr.x)
def simple_substep(self, x, sampler):
for sr in self.substep(x, sampler):
def simple_substep(self, x, sampler, ss=None):
for sr in self.substep(x, sampler, ss=ss):
if not sr.final:
raise RuntimeError("Unexpected non-final sampler result in substep!")
sr.noise_x(ss=fallback(ss, self.ss))
return sr
def merge_steps(self, x, result=None, *, noise=None, ss=None, denoised=True):
@@ -127,17 +104,6 @@ class MergeSubstepsSampler:
preview_mode = fallback(preview_mode, self.preview_mode)
return ss.callback(hi=mr, preview_mode=preview_mode)
def call_model(self, x, ss=None, sigma=None, **kwargs):
ss = fallback(ss, self.ss)
sigma = fallback(sigma, ss.sigma)
return ss.call_model(
x,
sigma,
ss=ss,
cfg_scale_override=self.cfg_scale_override,
require_uncond=self.require_uncond,
)
class SimpleSubstepsSampler(MergeSubstepsSampler):
name = "simple"
@@ -152,16 +118,26 @@ class SimpleSubstepsSampler(MergeSubstepsSampler):
def step(self, x):
ss, ssampler = self.ss, self.samplers[0]
ss.hist.push(self.call_model(x))
custom_noise = ssampler.options.get(
"custom_noise", self.options.get("custom_noise")
)
noise_sampler = ss.noise.make_caching_noise_sampler(
custom_noise,
1,
ss.sigma,
ss.sigma_next,
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
)
ssampler.noise_sampler = noise_sampler
ss.hist.push(ss.model(x, ss.sigma, ss=ss))
ss.refs = FilterRefs.from_ss(ss, have_current=True)
self.callback()
with StepSamplerContext(ssampler, ss) as ssampler:
sr = self.simple_substep(x, ssampler)
return self.merge_steps(sr.noise_x(ss=ss))
sr = self.simple_substep(x, ssampler)
return self.merge_steps(sr.x, noise=sr.get_noise(ss=ss))
class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
name = "supreme_avg"
class NormalMergeSubstepsSampler(MergeSubstepsSampler):
name = "normal"
def step(self, x):
ss = self.ss
@@ -172,21 +148,31 @@ class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
noise_total = 0.0
substep = 0
pbar = tqdm.tqdm(total=self.substeps, initial=1, disable=ss.disable_status)
ss.hist.push(self.call_model(x))
ss.hist.push(ss.model(x, ss.sigma, ss=ss))
ss.refs = FilterRefs.from_ss(ss, have_current=True)
self.callback()
for ssampler_ in self.samplers:
with StepSamplerContext(ssampler_, ss) as ssampler:
for subidx in range(ssampler.substeps):
pbar.set_description(f"{ssampler.name}: {substep + 1}/{substeps}")
sr = self.simple_substep(x, ssampler)
z_avg += renoise_weight * sr.x
if sr.noise_scale != 0 and ss.sigma_next != 0:
noise_total += renoise_weight * sr.noise_scale
noise += renoise_weight * sr.get_noise(ss=ss)
substep += 1
ss.substep = substep
pbar.update(1)
for ssampler in self.samplers:
custom_noise = ssampler.options.get(
"custom_noise", self.options.get("custom_noise")
)
noise_sampler = ss.noise.make_caching_noise_sampler(
custom_noise,
ssampler.max_noise_samples(),
ss.sigma,
ss.sigma_next,
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
)
ssampler.noise_sampler = noise_sampler
for subidx in range(ssampler.substeps):
pbar.set_description(f"{ssampler.name}: {substep + 1}/{substeps}")
sr = self.simple_substep(x, ssampler)
z_avg += renoise_weight * sr.x
if sr.noise_scale != 0 and ss.sigma_next != 0:
noise_total += renoise_weight * sr.noise_scale
noise += renoise_weight * sr.get_noise(ss=ss)
substep += 1
ss.substep = substep
pbar.update(1)
noise = ss.noise.scale_noise(
noise,
@@ -241,7 +227,7 @@ class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
# noise_sampler = ss.noise.make_caching_noise_sampler(
# custom_noise,
# ssampler.substeps
# + (0 if ss.sigma_next == 0 else ssampler.max_noise_samples),
# + (0 if ss.sigma_next == 0 else ssampler.max_noise_samples()),
# ss.sigma,
# ss.sigma_next,
# )
@@ -303,7 +289,7 @@ class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
# )
# noise_sampler = ss.noise.make_caching_noise_sampler(
# custom_noise,
# ssampler.max_noise_samples + ssampler.substeps,
# ssampler.max_noise_samples() + ssampler.substeps,
# ss.sigma,
# ss.sigma_next,
# )
@@ -346,7 +332,7 @@ class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
# final = merge_ss.sigma_next == 0
# noise_sampler = merge_ss.noise.make_caching_noise_sampler(
# msampler.options.get("custom_noise", self.options.get("custom_noise")),
# msampler.max_noise_samples + int(not final),
# msampler.max_noise_samples() + int(not final),
# merge_ss.sigma,
# merge_ss.sigma_next,
# )
@@ -407,24 +393,34 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler):
subss.main_sigmas = ss.sigmas
substep = 0
pbar = tqdm.tqdm(total=self.substeps, initial=0, disable=ss.disable_status)
for ssampler_ in self.samplers:
with StepSamplerContext(ssampler_, subss) as ssampler:
for subidx in range(ssampler.substeps):
subss.update(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)
for ssampler in self.samplers:
custom_noise = ssampler.options.get(
"custom_noise", self.options.get("custom_noise")
)
noise_sampler = ss.noise.make_caching_noise_sampler(
custom_noise,
ssampler.max_noise_samples(),
ss.sigma,
ss.sigma_next,
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
)
ssampler.noise_sampler = noise_sampler
for subidx in range(ssampler.substeps):
subss.update(substep, substep=substep)
pbar.set_description(
f"substep({ssampler.name}): {subss.sigma.item():.03} -> {subss.sigma_next.item():.03}"
)
subss.hist.push(subss.model(x, subss.sigma, ss=subss))
subss.refs = FilterRefs.from_ss(subss, have_current=True)
if substep == 0:
self.callback(ss=subss)
sr = self.simple_substep(x, ssampler, ss=subss)
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)
pbar.update(0)
return x
@@ -441,16 +437,10 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
super().__init__(ss, group, **kwargs)
self.overshoot_expand_steps = self.options.pop("overshoot_expand_steps", 1)
restart = self.options.pop("restart", {})
restart_custom_noise = self.options.get("restart_custom_noise")
if isinstance(restart_custom_noise, str):
restart_custom_noise = self.options.get(
f"restart_custom_noise_{restart_custom_noise}"
)
self.restart = Restart(
s_noise=restart.get("s_noise", 1.0),
custom_noise=restart_custom_noise,
custom_noise=self.options.pop("restart_custom_noise", None),
immiscible=restart.get("immiscible", False),
is_flow=ss.model.is_rectified_flow,
)
def make_schedule(self, ss):
@@ -480,285 +470,57 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
pbar = tqdm.tqdm(total=self.substeps, initial=0, disable=ss.disable_status)
max_idx = len(subss.sigmas) - 2
last_down = None
for ssampler_ in self.samplers:
with StepSamplerContext(ssampler_, subss) as ssampler:
for subidx in range(ssampler.substeps):
subss.update(subss.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:
ss.hist.push(subss.hcur)
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)
last_down = subss.sigma_next.item()
if subss.idx + substep >= max_idx:
break
if subss.idx >= max_idx:
for ssampler in self.samplers:
custom_noise = ssampler.options.get(
"custom_noise", self.options.get("custom_noise")
)
noise_sampler = ss.noise.make_caching_noise_sampler(
custom_noise,
ssampler.max_noise_samples(),
ss.sigma,
ss.sigma_next,
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
)
ssampler.noise_sampler = noise_sampler
for subidx in range(ssampler.substeps):
subss.update(subss.idx + substep, substep=substep)
pbar.set_description(
f"substep({ssampler.name}): {subss.sigma.item():.03} -> {subss.sigma_next.item():.03}"
)
subss.hist.push(subss.model(x, subss.sigma, ss=subss))
subss.refs = FilterRefs.from_ss(subss, have_current=True)
if substep == 0:
ss.hist.push(subss.hcur)
self.callback(ss=subss)
sr = self.simple_substep(x, ssampler, ss=subss)
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)
last_down = subss.sigma_next.item()
if subss.idx + substep >= max_idx:
break
if subss.idx >= max_idx:
break
if last_down is not None and last_down < ss.sigma_next:
x = self.restart.add_noise(
x,
sigma_from=last_down.item(),
sigma_to=ss.sigma_next.item(),
nsc=nsc,
refs=ss.refs,
in_place=True,
restart_ns = self.restart.get_noise_sampler(ss.noise)
x += ss.noise.scale_noise(
restart_ns(refs=ss.refs),
self.restart.get_noise_scale(last_down, ss.sigma_next),
)
pbar.update(0)
return x
class LookaheadMergeSubstepsSampler(MergeSubstepsSampler):
name = "lookahead"
def __init__(self, ss, group, **kwargs):
super().__init__(ss, group, **kwargs)
lookahead = self.options.pop("lookahead", {}).copy()
self.lookahead_eta = lookahead.pop("eta", 0.0)
self.lookahead_s_noise = lookahead.pop("s_noise", 1.0)
self.lookahead_dt_factor = lookahead.pop("dt_factor", 1.0)
immiscible = lookahead.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)
sigma_down, sigma_up = ss.get_ancestral_step(
eta=self.lookahead_eta, sigma=ss.sigma, sigma_next=ss.sigma_next
)
if sr.sigma_next == sigma_down:
return x
dt = (
torch.sqrt(1.0 + (ss.sigma - sigma_down) ** 2) * 0.05
+ (ss.sigma - sigma_down) * 0.95
) * self.lookahead_dt_factor
denoised = sr.denoised
d = (orig_x - denoised) / ss.sigma
x = orig_x + d * -dt
if sigma_down == 0 or sigma_up == 0:
return x
noise_sampler = ss.noise.make_caching_noise_sampler(
self.custom_noise,
1,
ss.sigma,
ss.sigma_next,
immiscible=fallback(self.immiscible, ss.noise.immiscible),
)
# FIXME: This sigma, sigma_next is probably wrong.
x += ss.noise.scale_noise(
noise_sampler(ss.sigma, ss.sigma_next, refs=ss.refs),
sigma_up * self.lookahead_s_noise,
)
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"
def __init__(self, ss, group, **kwargs):
super().__init__(ss, group, **kwargs)
dynamic = self.options.get("dynamic")
if dynamic is None:
raise ValueError(
"Dynamic group 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):
group_params = None
handlers = 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
group_params = dyn_params.eval(handlers)
if group_params is not None:
break
if group_params is None:
raise RuntimeError(
"Dynamic group could not find matching group: all expressions failed to return a result"
)
if not isinstance(group_params, dict):
raise TypeError(
f"Dynamic group expression must evaluate to a dict, got type {type(group_params)}"
)
if bool(group_params.get("dynamic_inherit")):
copy_keys = ("preview_mode",)
opts = {k: getattr(self, k) for k in copy_keys}
else:
opts = {}
opts |= {
k: v
for k, v in self.options.items()
if k.startswith("custom_noise") or k.startswith("restart_custom_noise")
}
opts |= group_params
# print("\n\nDYN GROUP OPTS", opts)
merge_method = opts.pop("merge_method", "simple").strip()
if merge_method == "default":
merge_method = "simple"
group_class = MERGE_SUBSTEPS_CLASSES.get(merge_method)
if group_class is None:
raise ValueError(f"Unknown merge method {merge_method} in dynamic group")
group = StepSamplerChain(
merge_method=merge_method, items=self.group.items, **opts
)
sampler = group_class(self.ss, group)
return sampler.step(x)
MERGE_SUBSTEPS_CLASSES = {
"default (simple)": SimpleSubstepsSampler,
"supreme_avg": SupremeAvgMergeSubstepsSampler,
"normal": NormalMergeSubstepsSampler,
"divide": DivideMergeSubstepsSampler,
"overshoot": OvershootMergeSubstepsSampler,
# "average": AverageMergeSubstepsSampler,
# "sample": SampleMergeSubstepsSampler,
# "sample_uncached": SampleUncachedMergeSubstepsSampler,
"simple": SimpleSubstepsSampler,
"lookahead": LookaheadMergeSubstepsSampler,
"pingpong": PingpongMergeSubstepsSampler,
"dynamic": DynamicMergeSubstepsSampler,
}
+16 -165
View File
@@ -1,6 +1,5 @@
from typing import NamedTuple
import torch
from comfy.k_diffusion.sampling import get_ancestral_step
from .filtering import FilterRefs
@@ -8,13 +7,6 @@ from .model import History
from .utils import fallback
class AncestralRatios(NamedTuple):
alpha_t: torch.Tensor
alpha_s: torch.Tensor
sigma_up: torch.Tensor
sigma_down: torch.Tensor
class Items:
def __init__(self, items=None):
self.items = [] if items is None else items
@@ -90,7 +82,6 @@ class StepSamplerGroups(CommonOptionsItems):
class SamplerState:
CLONE_KEYS = (
"cfg_scale_override",
"model",
"hist",
"extra_args",
@@ -132,7 +123,6 @@ class SamplerState:
s_noise=1.0,
disable_status=False,
history_size=4,
cfg_scale_override=None,
):
self.model = model
self.hist = History(max(1, history_size))
@@ -148,11 +138,6 @@ class SamplerState:
self.step = 0
self.substep = 0
self.total_steps = len(sigmas) - 1
self.cfg_scale_override = cfg_scale_override
self.is_flow = self.model.is_rectified_flow
self.offset_sigma = (
model.model_sampling.percent_to_sigma(1e-04) if self.is_flow else None
)
self.update(idx) # Sets idx, sigma_prev, sigma, sigma_down, refs
@property
@@ -167,14 +152,6 @@ 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
@@ -183,23 +160,6 @@ class SamplerState:
def d(self):
return self.hcur.d
# These two functions referenced from ComfyUI.
def sigma_to_half_log_snr(
self, *, sigma: torch.Tensor | None = None, idx: int | None = None
) -> torch.Tensor:
if sigma is None and idx is None:
sigma = self.sigma
else:
sigma = sigma if sigma is not None else self.sigmas[idx]
if not self.is_flow:
return sigma.log().neg_()
if sigma.max() >= 1.0:
sigma = sigma * 0.0 + self.offset_sigma
return sigma.logit().neg_()
def half_log_snr_to_sigma(self, half_log_snr: torch.Tensor) -> torch.Tensor:
return (torch.sigmoid if self.is_flow else torch.exp)(half_log_snr.neg())
def update(self, idx=None, step=None, substep=None):
idx = self.idx if idx is None else idx
self.idx = idx
@@ -214,109 +174,16 @@ class SamplerState:
self.substep = substep
self.refs = FilterRefs.from_ss(self)
def get_ancestral_step_ext(
self,
*,
sigma: torch.Tensor | None = None,
sigma_next: torch.Tensor | None = None,
eta: float = 1.0,
retry_increment: int = 0,
):
sigma = fallback(sigma, self.sigma)
sigma_next = fallback(sigma_next, self.sigma_next)
sigma_empty = sigma_next * 0.0
def get_noeta_ratios():
return AncestralRatios(
alpha_t=sigma_empty + 1.0,
alpha_s=sigma_empty + 1.0,
sigma_up=sigma_empty.clone(),
sigma_down=sigma_next.clone(),
def get_ancestral_step(self, eta=1.0, sigma=None, sigma_next=None):
sigma = self.sigma if sigma is None else sigma
sigma_next = self.sigma_next if sigma_next is None else sigma_next
sd, su = (
v if isinstance(v, torch.Tensor) else sigma.new_full((1,), v)
for v in get_ancestral_step(
sigma, sigma_next, eta=eta if sigma_next != 0 else 0
)
if eta <= 0 or sigma_next.max().item() <= 1e-08:
return get_noeta_ratios()
orig_dtype = sigma.dtype
sigma = sigma.to(dtype=torch.float64)
sigma_next = sigma_next.to(dtype=torch.float64)
alpha_s = sigma * self.sigma_to_half_log_snr(sigma=sigma).exp()
alpha_t = sigma_next * self.sigma_to_half_log_snr(sigma=sigma_next).exp()
adj_sigma = sigma / alpha_s
adj_sigma_next = sigma_next / alpha_t
sd = su = None
while eta > 0:
sd, su = (
v if isinstance(v, torch.Tensor) else sigma.new_full((1,), v)
for v in get_ancestral_step(adj_sigma, adj_sigma_next, eta=eta)
)
if sd > 0 and su > 0:
break
else:
sd = su = None
if retry_increment <= 0:
break
# print(f"\nETA {eta} failed, retrying with {eta - retry_increment}")
eta -= retry_increment
if sd is None or su is None:
return get_noeta_ratios()
sd = alpha_t * sd
return AncestralRatios(
alpha_t=alpha_t.to(dtype=orig_dtype),
alpha_s=alpha_s.to(dtype=orig_dtype),
sigma_up=su.to(dtype=orig_dtype),
sigma_down=sd.to(dtype=orig_dtype),
)
def get_ancestral_step(
self, eta=1.0, sigma=None, sigma_next=None, retry_increment=0
):
if self.model.is_rectified_flow:
return self.get_ancestral_step_rf(
eta=eta,
sigma=sigma,
sigma_next=sigma_next,
retry_increment=retry_increment,
)
sigma = fallback(sigma, self.sigma)
sigma_next = fallback(sigma_next, self.sigma_next)
if eta <= 0 or sigma_next <= 0:
return sigma_next, sigma_next.new_zeros(1)
while eta > 0:
sd, su = (
v if isinstance(v, torch.Tensor) else sigma.new_full((1,), v)
for v in get_ancestral_step(
sigma, sigma_next, eta=eta if sigma_next != 0 else 0
)
)
if sd > 0 and su > 0:
return sd, su
if retry_increment <= 0:
break
# print(f"\nETA {eta} failed, retrying with {eta - retry_increment}")
eta -= retry_increment
return sigma_next, sigma_next.new_zeros(1)
# Referenced from Comfy dpmpp_2s_ancestral_RF
def get_ancestral_step_rf(
self, eta=1.0, sigma=None, sigma_next=None, retry_increment=0
):
sigma = fallback(sigma, self.sigma)
sigma_next = fallback(sigma_next, self.sigma_next)
if eta <= 0 or sigma_next <= 0:
return sigma_next, sigma_next.new_zeros(1)
while eta > 0:
sigma_down = sigma_next * (1 + (sigma_next / sigma - 1) * eta)
alpha_ip1, alpha_down = 1 - sigma_next, 1 - sigma_down
sigma_up = (
sigma_next**2 - sigma_down**2 * alpha_ip1**2 / alpha_down**2
) ** 0.5
if sigma_down > 0 and sigma_up > 0:
return sigma_down, sigma_up
if retry_increment <= 0:
break
eta -= retry_increment
return sigma_next, sigma_next.new_zeros(1)
# print(f"\nRF ancestral: down={sigma_down}, up={sigma_up}")
return sd, su
def clone_edit(self, **kwargs):
obj = self.__class__.__new__(self.__class__)
@@ -335,32 +202,16 @@ class SamplerState:
preview = fallback(hi.denoised_uncond, hi.denoised)
elif preview_mode == "raw":
preview = hi.x
elif (
preview_mode == "diff"
and hi.denoised_uncond is not None
and hi.denoised_cond is not None
):
preview = (
hi.denoised_uncond * 0.25 + (hi.denoised_uncond - hi.denoised_cond) * 16
)
elif preview_mode == "noisy":
preview = (hi.x - hi.denoised) * 0.1 + hi.denoised
else:
preview = hi.denoised
return self.callback_(
{
"x": hi.x,
"i": self.step,
"sigma": hi.sigma,
"sigma_hat": hi.sigma,
"denoised": preview,
}
)
return self.callback_({
"x": hi.x,
"i": self.step,
"sigma": hi.sigma,
"sigma_hat": hi.sigma,
"denoised": preview,
})
def reset(self):
self.hist.reset()
self.denoised = None
def call_model(self, *args, **kwargs):
cfg_scale_override = kwargs.pop("cfg_scale_override", self.cfg_scale_override)
return self.model(*args, cfg_scale_override=cfg_scale_override, **kwargs)
-386
View File
@@ -1,386 +0,0 @@
import torch
import triton
import triton.language as tl
@triton.autotune(
configs=[
triton.Config({"num_warps": 4, "num_stages": 2}, num_warps=4, num_stages=2),
triton.Config({"num_warps": 8, "num_stages": 2}, num_warps=8, num_stages=2),
triton.Config({"num_warps": 4, "num_stages": 3}, num_warps=4, num_stages=3),
triton.Config({"num_warps": 8, "num_stages": 3}, num_warps=8, num_stages=3),
],
key=[
"B",
"R",
"C",
"BLOCK_SIZE",
], # Retune if matrix dimensions change significantly
)
@triton.jit
def auction_lap_kernel(
cost_ptr,
assign_ptr,
stride_b,
stride_r,
stride_c,
stride_assign_b,
stride_assign_r,
B: tl.constexpr,
R: tl.constexpr,
C: tl.constexpr,
epsilon,
max_iter,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
cost_base = cost_ptr + pid * stride_b
assign_base = assign_ptr + pid * stride_assign_b
offs = tl.arange(0, BLOCK_SIZE)
col_mask = offs < C
# Prices and Owners in SRAM/Registers
prices = tl.zeros([BLOCK_SIZE], dtype=tl.float32)
owners = tl.full([BLOCK_SIZE], -1, dtype=tl.int32)
row_to_col = tl.full([BLOCK_SIZE], -1, dtype=tl.int32)
iter_idx = 0
unassigned_count = R
loop_continue = tl.full([], 1, dtype=tl.int1)
# Loop condition:
# 1. unassigned_count > 0: Logic handled inside, but we need a break mechanism
# 2. iter_idx < max_iter: Safety break
# 3. loop_continue: Did we make progress last time?
while unassigned_count > 0 and iter_idx < max_iter and loop_continue:
# Reset progress flag
# loop_continue &= False
loop_continue = tl.full([], 0, dtype=tl.int1)
# Gauss-Seidel pass over all rows
for i in tl.range(0, R):
# Check if row i is unassigned
curr_c = tl.sum(tl.where(offs == i, row_to_col, 0))
if curr_c == -1:
# Load costs
row_cost_ptr = cost_base + i * stride_r + offs
row_costs = tl.load(row_cost_ptr, mask=col_mask, other=-torch.inf)
# Net Value
values = row_costs - prices
# Find Best
best_val, best_idx = tl.max(values, axis=0, return_indices=True)
# CRITICAL: Only proceed if this is a valid edge (not -inf)
if best_val > -torch.inf:
# We have a valid move, so we continue the outer loop
loop_continue = tl.full([], 1, dtype=tl.int1)
# Find Second Best
mask_not_best = (offs != best_idx) & col_mask
vals_no_best = tl.where(mask_not_best, values, -torch.inf)
second_best_val = tl.max(vals_no_best, axis=0)
# Compute Bid
bid = best_val - second_best_val + epsilon
# Update Price
prices = tl.where(offs == best_idx, prices + bid, prices)
# Update Owners
prev_owner = tl.sum(tl.where(offs == best_idx, owners, 0))
if prev_owner != -1:
# Kick out previous owner
row_to_col = tl.where(offs == prev_owner, -1, row_to_col)
unassigned_count += 1
# Assign to current row
owners = tl.where(offs == best_idx, i, owners)
row_to_col = tl.where(offs == i, best_idx, row_to_col)
unassigned_count -= 1
iter_idx += 1
# Store Result
store_offs = tl.arange(0, BLOCK_SIZE)
store_mask = store_offs < R
tl.store(assign_base + store_offs, row_to_col, mask=store_mask)
# -----------------------------------------------------------------------------
# Python Helpers
# -----------------------------------------------------------------------------
def rescale_simple(
t: torch.Tensor,
target_min: float = 0.0,
target_max: float = 1.0,
*,
start_dim: int = 1,
eps: float = 1e-07,
) -> torch.Tensor:
width = target_max - target_min
if width == 0.0:
return torch.zeros_like(t)
orig_shape = t.shape
t = t.flatten(start_dim=start_dim)
min_val, max_val = t.aminmax(dim=-1, keepdim=True)
normalized = t - min_val
normalized /= (max_val - min_val).add_(eps)
normalized *= width
if target_min != 0.0:
normalized += target_min
return normalized.clamp_(target_min, target_max).reshape(orig_shape)
def _greedy_fill_missing(assignments: torch.Tensor, C: int) -> None:
"""
Fills unassigned rows (-1) in the assignments tensor with available columns.
This acts as a fallback when the Auction algorithm hits max_iter without
full convergence.
Args:
assignments: Tensor of shape (B, R) containing col indices or -1.
C: Total number of columns available.
"""
# Identify which batch items have unassigned rows
# This is usually a very small subset (e.g., < 1% of the batch)
problem_batches = (assignments == -1).any(dim=1).nonzero().flatten()
if problem_batches.numel() == 0:
return
device = assignments.device
# Iterate only over the problematic batch items
# (Looping is acceptable here as B_subset is typically tiny)
for b_idx in problem_batches:
# 1. Find which rows are missing an assignment
row_mask = assignments[b_idx] == -1
missing_rows = row_mask.nonzero().flatten()
n_needed = missing_rows.shape[0]
# 2. Find which columns are already used
used_cols = assignments[b_idx][~row_mask]
# 3. Find free columns (Set difference: All - Used)
# Create a boolean mask of all columns, then mark used ones as False
# efficient on GPU for mid-sized C
col_mask = torch.ones(C, device=device, dtype=torch.bool)
col_mask[used_cols.long()] = False
free_cols = col_mask.nonzero().flatten()
# 4. Assign the first N free columns to the N missing rows
# Since R <= C in this context (due to transpose logic in wrapper),
# free_cols.numel() is guaranteed to be >= n_needed.
assignments[b_idx, missing_rows] = free_cols[:n_needed].to(assignments.dtype)
def batch_linear_assignment(
cost_matrix: torch.Tensor,
*,
maximize: bool = False,
max_iter: int | None = None,
fill_missing: bool = True,
rescale_costs: tuple[float, float] | None = (0.0, 1.0),
invert_costs_mode: bool = True,
eps: float = 1e-3,
):
if cost_matrix.ndim != 3:
raise ValueError("Cost matrix must be (B, R, C)")
if not cost_matrix.is_cuda:
raise ValueError("Cost matrix must be a CUDA tensor")
if not cost_matrix.is_contiguous():
raise ValueError("Cost matrix must be contiguous")
B, R, C = cost_matrix.shape
device = cost_matrix.device
# 1. Handle Rectangular Matrices
# The Auction algorithm assigns Rows -> Cols.
# It naturally handles R <= C (finding best col for every row).
# If R > C, we must transpose to match Cols -> Rows, then invert the result.
if R > C:
transposed = True
cost_matrix = cost_matrix.mT.contiguous()
# Swap R and C for the kernel execution
R, C = C, R
else:
transposed = False
if rescale_costs is not None:
cost_matrix = rescale_simple(cost_matrix, *rescale_costs)
# Note: We use float32 for atomic compatibility and speed
cost_matrix = cost_matrix.to(torch.float32, copy=rescale_costs is None)
if not maximize:
# Maximize (Value - Price) -> Minimize Cost
cost_matrix = cost_matrix.neg_()
if invert_costs_mode and rescale_costs is not None:
cost_matrix += sum(rescale_costs)
assignments = torch.full(
(B, R),
-1,
device=device,
dtype=torch.int32,
)
max_dim = max(R, C)
BLOCK_SIZE = max(32, triton.next_power_of_2(max_dim))
# Safety limit
max_iter = max_iter if max_iter is not None else int(max(2000, R * C))
grid = (B,)
auction_lap_kernel[grid](
cost_matrix,
assignments,
cost_matrix.stride(0),
cost_matrix.stride(1),
cost_matrix.stride(2),
assignments.stride(0),
assignments.stride(1),
B,
R,
C,
eps,
max_iter,
BLOCK_SIZE=BLOCK_SIZE,
)
if fill_missing:
_greedy_fill_missing(assignments, C)
assignments = assignments.long()
if not transposed:
return assignments
# 2. Post-process Rectangular Results
# We computed Col -> Row. We need Row -> Col.
# assignments shape is currently (B, Original_Cols)
# We want output shape (B, Original_Rows)
real_rows = C # C is the 'large' dimension (Original Rows)
output = torch.full((B, real_rows), -1, device=device, dtype=torch.long)
# Create indices for the scatter source
# We want: output[row_idx] = col_idx
# Currently we have: assignments[col_idx] = row_idx
src_col_indices = torch.arange(R, device=device).unsqueeze(0).expand(B, R)
# We use scatter. index=assignments (the rows), src=col_indices
# To handle -1s in assignments, we clamp to 0 and then mask the result
safe_assigns = assignments.clamp(min=0)
output.scatter_(1, safe_assigns, src_col_indices)
# Cleanup: Any row that wasn't targeted by the scatter should be -1
# The scatter might have written to index 0 if assignment was -1
# Re-verify logic:
for b in range(B):
valid_mask = assignments[b] >= 0
# Reset output
output[b].fill_(-1)
# Only write valid mappings
# output[b, row_id] = col_id
output[b, assignments[b, valid_mask]] = src_col_indices[b, valid_mask]
return output
def assignments_to_indices(
assignments: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Converts a dense assignment tensor (from Triton/Hungraian) to
batched SciPy-style indices.
Args:
assignments (torch.Tensor): Shape (B, R). Values are col indices or -1.
Returns:
row_ind (torch.Tensor): Shape (B, K) where K = min(R, C).
col_ind (torch.Tensor): Shape (B, K).
"""
B, R = assignments.shape
device = assignments.device
# 1. Create a mask of valid assignments (values >= 0)
# In a rectangular assignment, the number of valid matches
# is always min(Rows, Cols).
mask = assignments >= 0
# 2. Extract Column Indices
# We select the values from the assignment tensor that are valid.
# We reshape to (B, -1) to preserve the batch dimension.
col_ind = assignments[mask].view(B, -1)
# 3. Extract Row Indices
# We need a grid of row indices [0, 1, 2, ... R-1] repeated B times
row_grid = (
torch.arange(R, device=device, dtype=assignments.dtype)
.unsqueeze(0)
.expand(B, R)
)
row_ind = row_grid[mask].view(B, -1)
return row_ind, col_ind
def batch_linear_assignment_shuffled(
cost_matrix: torch.Tensor,
*args,
**kwargs: dict,
) -> torch.Tensor:
generator = kwargs.pop("generator", None)
# cost_matrix: [B, R, C]
B, R = cost_matrix.shape[:2]
# 1. Generate a random permutation for the rows
# We use one perm for the whole batch for efficiency,
# or you can do it per-batch-item if B is small and quality is critical.
# Here we shuffle all rows commonly.
perm = torch.randperm(
R,
device=cost_matrix.device,
generator=generator,
)
# 2. Shuffle the input (Row dimension is dim 1)
# This creates a shuffled view/copy of the cost matrix
shuffled_cost = cost_matrix[:, perm, :]
# 3. Run the Solver
shuffled_assignments = batch_linear_assignment(
shuffled_cost,
*args,
**kwargs,
) # Returns [B, R]
# 4. Un-shuffle the results
# We need to map the results back to their original row positions.
# shuffled_assignments[b, i] corresponds to the row 'perm[i]'
# We want final_assignments[b, perm[i]] = shuffled_assignments[b, i]
# Create the inverse permutation or just scatter back
final_assignments = torch.empty_like(shuffled_assignments)
# Expand perm for the batch: [B, R]
batch_perm = perm.unsqueeze(0).expand(B, R)
# Scatter the results back to original positions
# dim=1, index=batch_perm, src=shuffled_assignments
final_assignments.scatter_(1, batch_perm, shuffled_assignments)
return final_assignments
-287
View File
@@ -1,287 +0,0 @@
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",
))
+61 -103
View File
@@ -1,119 +1,77 @@
from __future__ import annotations
import contextlib
import torch
from comfy.k_diffusion.sampling import to_d
F = torch.nn.functional
from . import latent
# def scale_noise_(
# noise,
# factor=1.0,
# *,
# normalized=True,
# normalize_dims=(-3, -2, -1),
# ):
# if not normalized or noise.numel() == 0:
# return noise.mul_(factor) if factor != 1 else noise
# mean, std = (
# noise.mean(dim=normalize_dims, keepdim=True),
# noise.std(dim=normalize_dims, keepdim=True),
# )
# return latent.normalize_to_scale(
# noise.sub_(mean).div_(std).clamp(-1, 1), -1.0, 1.0, dim=normalize_dims
# ).mul_(factor)
# def scale_noise(
# noise,
# factor=1.0,
# *,
# normalized=True,
# normalize_dims=(-3, -2, -1),
# ):
# if not normalized or noise.numel() == 0:
# return noise * factor if factor != 1 else noise
# mean, std = (
# noise.mean(dim=normalize_dims, keepdim=True),
# noise.std(dim=normalize_dims, keepdim=True),
# )
# return (noise - mean).div_(std).mul_(factor)
def scale_noise(
noise: torch.Tensor,
factor: float = 1.0,
noise,
factor=1.0,
*,
normalized: bool = True,
normalize_dims: tuple[int, ...] | None = None,
eps: float | None = None,
) -> torch.Tensor:
if factor == 0:
return torch.zeros_like(noise)
normalized=True,
normalize_dims=(-3, -2, -1),
):
if not normalized or noise.numel() == 0:
return noise * factor if factor != 1 else noise
if eps is None:
eps = torch.finfo(noise.dtype).eps * 1.25
if normalize_dims is None:
normalize_dims = tuple(
range(
max(0, min(1, noise.ndim - 1)),
noise.ndim,
)
)
std, mean = torch.std_mean(noise, dim=normalize_dims, keepdim=True)
noise = noise - mean
if factor != 1:
std /= factor
return noise.div_(std.clamp_min_(eps) if factor >= 0 else std.clamp_max_(-eps))
noise = noise / noise.std(dim=normalize_dims, keepdim=True)
return noise.sub_(noise.mean(dim=normalize_dims, keepdim=True)).mul_(factor)
def range_wrap(
x: torch.Tensor,
min_val: float | torch.Tensor,
max_val: float | torch.Tensor,
) -> torch.Tensor:
return min_val + (x - min_val).remainder_(max_val - min_val)
def softplus_soft_clamp(
t: torch.Tensor,
min_val: torch.Tensor | float = 0.0,
max_val: torch.Tensor | float = 1.0,
*,
# We define stiffness as a multiplier (beta) for the softplus function.
# Higher stiffness = sharper transition.
stiffness: float = 1.0,
safe: bool = True,
) -> torch.Tensor:
if isinstance(min_val, (float, int)):
min_val = t.new_tensor(min_val)
if isinstance(max_val, (float, int)):
max_val = t.new_tensor(max_val)
if stiffness < 1e-04:
return t.clamp(min_val, max_val)
# Calculate how much we are exceeding the Max
# softplus(beta * x) / beta
upper_overshoot = F.softplus((t - max_val).mul_(stiffness)).div_(-stiffness)
# Calculate how much we are falling short of the Min
lower_undershoot = F.softplus((min_val - t).mul_(stiffness)).div_(stiffness)
# Apply corrections:
# Original - (Amount over max) + (Amount under min)
t = upper_overshoot.add_(t).add_(lower_undershoot)
if safe:
t = t.clamp(min_val, max_val)
return t
def flip_tensor_range(
x: torch.Tensor,
*,
min_neg: torch.Tensor | None = None,
max_pos: torch.Tensor | None = None,
return_ranges: bool = False,
dim: int = -1,
eps: float | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
if eps is None:
eps = torch.finfo(x.dtype).eps * 1.25
# 1. Use the provided maximum positive values, or calculate them dynamically
if max_pos is None:
max_pos = (
torch.clamp_min(x, 0.0).max(dim=dim, keepdim=True).values.clamp_min_(eps)
)
# 2. Use the provided minimum negative values, or calculate them dynamically
if min_neg is None:
min_neg = (
torch.clamp_max(x, 0.0).min(dim=dim, keepdim=True).values.clamp_max_(-eps)
)
# 3. Separate positive and negative elements
is_pos = x >= 0
# 4. Flip positive side: [0, max_pos] -> [eps, max_pos + eps]
x_pos = x.clamp_min(eps)
flipped_pos = (max_pos + eps) - x_pos
# 5. Flip negative side: [min_neg, 0] -> [min_neg - eps, -eps]
x_neg = x.clamp_max(-eps)
flipped_neg = (min_neg - eps) - x_neg
# 6. Recombine the domains
result = torch.where(is_pos, flipped_pos, flipped_neg)
return (result, max_pos, min_neg) if return_ranges else result
# def scale_noise(
# noise,
# factor=1.0,
# *,
# normalized=True,
# normalize_dims=(-3, -2, -1),
# ):
# if not normalized or noise.numel() == 0:
# return noise.mul_(factor) if factor != 1 else noise
# n = (
# torch.nn.LayerNorm(noise.shape[1:])
# if normalize_dims == (-3, -2, -1)
# else torch.nn.InstanceNorm2d(noise.shape[1])
# ).to(noise)
# return n(noise) * factor
# return latent.normalize_to_scale(
# n(noise).clamp_(-1, 1), -1, 1, dim=normalize_dims
# ).mul_(factor)
def find_first_unsorted(tensor, desc=True):
-238
View File
@@ -1,238 +0,0 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Callable
import torch
from .utils import fallback
if TYPE_CHECKING:
from collections.abc import Sequence
try:
import pytorch_wavelets as ptwav
import pywt
HAVE_WAVELETS = True
except ImportError:
ptwav = None
pywt = None
HAVE_WAVELETS = False
class Wavelet:
DEFAULT_MODE = "symmetric"
DEFAULT_LEVEL = 3
DEFAULT_WAVE = "db4"
DEFAULT_USE_1D_DWT = False
DEFAULT_USE_DTCWT = False
DEFAULT_QSHIFT = "qshift_a"
DEFAULT_BIORT = "near_sym_a"
def __init__(
self,
*,
wave: str = DEFAULT_WAVE,
level: int = DEFAULT_LEVEL,
mode: str = DEFAULT_MODE,
use_1d_dwt: bool = DEFAULT_USE_1D_DWT,
use_dtcwt: bool = DEFAULT_USE_DTCWT,
biort: str = DEFAULT_BIORT,
qshift: str = DEFAULT_QSHIFT,
inv_wave: str | None = None,
inv_mode: str | None = None,
inv_biort: str | None = None,
inv_qshift=None,
device: str | torch.device | None = None,
):
if not HAVE_WAVELETS:
raise RuntimeError(
"Wavelet noise requires the pytorch_wavelets package to be installed in your Python environment",
)
inv_wave = fallback(inv_wave, wave)
inv_mode = fallback(inv_mode, mode)
inv_biort = fallback(inv_biort, biort)
inv_qshift = fallback(inv_qshift, qshift)
if use_dtcwt:
fwdfun, invfun = ptwav.DTCWTForward, ptwav.DTCWTInverse
elif use_1d_dwt:
fwdfun, invfun = ptwav.DWT1DForward, ptwav.DWT1DInverse
else:
fwdfun, invfun = ptwav.DWTForward, ptwav.DWTInverse
if use_dtcwt:
self._wavelet_forward = fwdfun(
J=level,
mode=mode,
biort=biort,
qshift=qshift,
)
self._wavelet_inverse = invfun(
mode=inv_mode,
biort=inv_biort,
qshift=inv_qshift,
)
else:
self._wavelet_forward = fwdfun(J=level, wave=wave, mode=mode)
self._wavelet_inverse = invfun(wave=inv_wave, mode=inv_mode)
if device is not None:
self._wavelet_forward = self._wavelet_forward.to(device=device)
self._wavelet_inverse = self._wavelet_inverse.to(device=device)
def forward(
self,
t: torch.Tensor,
*,
forward_function: Callable | None = None,
) -> tuple[torch.Tensor, tuple]:
return fallback(forward_function, self._wavelet_forward)(t)
def inverse(
self,
yl: torch.Tensor,
yh: tuple,
*,
inverse_function: Callable | None = None,
two_step_inverse: bool = False,
) -> torch.Tensor:
inverse_function = fallback(inverse_function, self._wavelet_inverse)
if not two_step_inverse:
return inverse_function((yl, yh))
result = inverse_function((torch.zeros_like(yl), yh))
result += inverse_function((
yl,
tuple(torch.zeros_like(yh_band) for yh_band in yh),
))
return result
def to(self, *args: list, copy: bool = False, **kwargs: dict) -> Wavelet:
o = Wavelet.__new__(Wavelet) if copy else self
o._wavelet_forward = self._wavelet_forward.to(*args, **kwargs) # noqa: SLF001
o._wavelet_inverse = self._wavelet_inverse.to(*args, **kwargs) # noqa: SLF001
return o
@staticmethod
def wavelist() -> tuple:
return tuple(pywt.wavelist()) if HAVE_WAVELETS else ()
@staticmethod
def biortlist() -> tuple:
return (
("near_sym_a", "near_sym_b", "antonini", "legall") if HAVE_WAVELETS else ()
)
@staticmethod
def qshiftlist() -> tuple:
return (
("qshift_a", "qshift_b", "qshift_c", "qshift_d", "qshift_06")
if HAVE_WAVELETS
else ()
)
@staticmethod
def modelist() -> tuple:
return (
(
"symmetric",
"zero",
"reflect",
"replicate",
"periodization",
"periodic",
"constant",
)
if HAVE_WAVELETS
else ()
)
def expand_yh_scales(
yh: Sequence,
*,
yh_scales: float | Sequence = 1.0,
) -> float | tuple:
yhlen = len(yh)
yh_shape = yh[0].shape
# Doesn't make sense to target orientations for 1D DWD (3D here).
olen = yh_shape[2] if len(yh_shape) > 3 else 1
# print(f"\nSIZES: yhlen={yhlen}, olen={olen}, yh_shape={yh[0].shape}")
if isinstance(yh_scales, (float, int)):
return ((float(yh_scales),) * olen,) * yhlen
otemplate = (1.0,) * olen
yh_scales = tuple(
(float(band),) * olen
if isinstance(band, (float, int))
else (
(
*(float(i) for i in band[:olen]),
*otemplate[: olen - len(band[:olen])],
)
if isinstance(band, (tuple, list))
else band
)
for band in yh_scales
)
if "fill" in yh_scales:
fillidx = yh_scales.index("fill")
if "fill" in yh_scales[fillidx + 1 :]:
raise ValueError("Only one fill allowed.")
if fillidx == 0 or len(yh_scales) < 2:
raise ValueError(
"Invalid fill value, cannot be in the first position or the only item.",
)
yhslen = len(yh_scales)
if yhslen - 1 < yhlen:
# Need to pad.
fill = (yh_scales[fillidx - 1],) * (yhlen - (len(yh_scales) - 1))
yh_scales = (*yh_scales[:fillidx], *fill, *yh_scales[fillidx + 1 :])
else:
# Just remove the "fill".
yh_scales = (*yh_scales[:fillidx], *yh_scales[fillidx + 1 :])
return yh_scales[:yhlen]
def wavelet_scaling(
yl: torch.Tensor,
yh: Sequence,
yl_scale: float | torch.Tensor,
yh_scales: float | Sequence | None,
*,
in_place: bool = False,
) -> tuple:
if not in_place:
yl = yl.clone()
yh = tuple(yhband.clone() for yhband in yh)
if yl_scale != 1.0:
yl *= yl_scale
yh_scales = expand_yh_scales(
yh,
yh_scales=yh_scales if yh_scales is not None else 1.0,
)
for hscale, ht in zip(yh_scales, yh):
if isinstance(hscale, (int, float)):
ht *= hscale # noqa: PLW2901
continue
for lidx in range(min(ht.shape[2], len(hscale))):
ht[:, :, lidx] *= hscale[lidx]
return (yl, yh)
def wavelet_blend(
a: tuple,
b: tuple,
*,
yl_factor: torch.Tensor | float,
blend_function: Callable,
yh_factor: torch.Tensor | float | None = None,
yh_blend_function: Callable | None = None,
) -> tuple:
if not isinstance(yl_factor, torch.Tensor):
yl_factor = a[0].new_full((1,), yl_factor)
if yh_factor is None:
yh_factor = yl_factor
elif not isinstance(yh_factor, torch.Tensor):
yh_factor = a[0].new_full((1,), yh_factor)
yh_blend_function = fallback(yh_blend_function, blend_function)
return (
blend_function(a[0], b[0], yl_factor),
tuple(yh_blend_function(ta, tb, yh_factor) for ta, tb in zip(a[1], b[1])),
)