Compare commits
20
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ee37473eac | ||
|
|
3d7193acd0 | ||
|
|
553127a766 | ||
|
|
0caf623d05 | ||
|
|
6a093ef6a0 | ||
|
|
222bd00991 | ||
|
|
2791278704 | ||
|
|
28b40df8b1 | ||
|
|
fd7ffc4e54 | ||
|
|
187c0cbabc | ||
|
|
e7db5e9213 | ||
|
|
361da0f32b | ||
|
|
ac194dc18d | ||
|
|
ef417b2e74 | ||
|
|
899fc19859 | ||
|
|
edad6fb0eb | ||
|
|
dcef1ea159 | ||
|
|
f9443ea74f | ||
|
|
4940e49196 | ||
|
|
9258cb2d90 |
@@ -4,40 +4,32 @@ Experimental and mathematically unsound (but fun!) sampling for [ComfyUI](https:
|
||||
|
||||
**Status**: In flux, may be useful but likely to change/break workflows frequently. Mainly for advanced users.
|
||||
|
||||
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
|
||||
@@ -57,15 +49,6 @@ You may use filters and expressions in the text parameter input. See:
|
||||
* [Filters](docs/filter.md)
|
||||
* [Expressions](docs/expression.md)
|
||||
|
||||
## Integration
|
||||
|
||||
* [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) - allows access to many more blend and scaling modes as well as some extra features.
|
||||
* [ComfyUI-sonar](https://github.com/blepping/ComfyUI-sonar) - allows access to many more noise types as well as the Power Filter feature.
|
||||
* [ComfyUi_NNLatentUpscale](https://github.com/Ttl/ComfyUi_NNLatentUpscale) - allows access to the `t_scale_nnlatentupscale` function in expressions.
|
||||
|
||||
If you're going to use OCS, I strongly recommend also installing `ComfyUI-bleh` and `ComfyUI-sonar` as they increase the functionality a lot.
|
||||
|
||||
|
||||
## Nodes
|
||||
|
||||
### `OCS Sampler`
|
||||
@@ -98,9 +81,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 +90,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,14 +113,14 @@ 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
|
||||
|
||||
# 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.
|
||||
cache_reset_interval: 9999
|
||||
# if using Brownian you will generally want to reset each step.
|
||||
cache_reset_interval: 1
|
||||
|
||||
# Immiscible noise processing, see: https://arxiv.org/abs/2406.12303
|
||||
immiscible:
|
||||
@@ -187,31 +163,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
|
||||
```
|
||||
|
||||
@@ -226,8 +201,6 @@ noise:
|
||||
|
||||
Then the rest of the parameters will use the defaults shown above.
|
||||
|
||||
***
|
||||
|
||||
### `OCS Group`
|
||||
|
||||
Defines a group of substeps.
|
||||
@@ -239,25 +212,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 +243,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
|
||||
|
||||
@@ -299,15 +261,6 @@ eta: 1.0
|
||||
# Reversible ETA (used for reversible samplers). May not do anything currently.
|
||||
reta: 1.0
|
||||
|
||||
# Sets the type of preview used for sampling in this group. One of:
|
||||
# denoised: The default, shows the model prediction (takes positive and negative prompt into account).
|
||||
# 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.
|
||||
when: null
|
||||
|
||||
@@ -323,26 +276,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
|
||||
@@ -350,52 +283,39 @@ post_filter: null
|
||||
|
||||
</details>
|
||||
|
||||
***
|
||||
|
||||
### `OCS Substeps`
|
||||
|
||||
#### Step Methods (Samplers)
|
||||
|
||||
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 +323,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 +375,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 +397,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 +414,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 +505,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 +517,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 +549,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.
|
||||
@@ -846,75 +582,3 @@ The example above means:
|
||||
3. Go to the second item (after 2 steps, jump back one step).
|
||||
|
||||
The node `start_step` parameter is effectively the same as `[start_step, 0]` as a schedule item.
|
||||
|
||||
***
|
||||
|
||||
### `OCSNoise to SONAR_CUSTOM_NOISE`
|
||||
|
||||
Adapter that enables using OCS noise generators with nodes that accept `SONAR_CUSTOM_NOISE`.
|
||||
|
||||
Most built-in OCS nodes will accept either type currently.
|
||||
|
||||
***
|
||||
|
||||
### `OCSNoise PerlinSimple`
|
||||
|
||||
Generates 2D or 3D Perlin noise with many tuneable parameters. Can be plugged in to samplers for ancestral or SDE sampling. For initial noise or img2img workflows, use the `NoisyLatentLike` node from `ComfyUI-sonar` (see [Integration](#integration)).
|
||||
|
||||
3D Perlin noise works by taking a slice in the depth dimension each time the noise sampler is called.
|
||||
|
||||
For more tuneable parameters, see the `OCSNoise PerlinAdvanced` node.
|
||||
|
||||
**Note**: The shape of the latent must be a multiple of `lacunarity ** (octaves - 1) * res` (`**` indicates raising something to a power). Most latent types will have one latent pixel equaling eight normal pixels - i.e. if your image is 512x512, the latent would be 64x64.
|
||||
|
||||
#### Node Parameters
|
||||
|
||||
* `depth`: When non-zero, 3D perlin noise will be generated.
|
||||
* `detail_level`: Controls the detail level of the noise when `break_pattern` is non-zero. No effect when using 100% raw Perlin noise.
|
||||
* `octaves`: Generally controls the detail level of the noise. Each octave involves generating a layer of noise so there is a performance cost to increasing octaves.
|
||||
* `persistence`: Controls how rough the generated noise is. Lower values will result in smoother noise, higher values will look more like Gaussian noise. Comma-separated list, multiple items will apply to octaves in sequence.
|
||||
* `lacunarity`: Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.
|
||||
* `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`
|
||||
|
||||
Generates 2D or 3D Perlin noise with many tuneable parameters. Can be plugged in to samplers for ancestral or SDE sampling. For initial noise or img2img workflows, use the `NoisyLatentLike` node from `ComfyUI-sonar` (see [Integration](#integration)).
|
||||
|
||||
3D Perlin noise works by taking a slice in the depth dimension each time the noise sampler is called.
|
||||
|
||||
**Note**: The shape of the latent in the relevant dimension _including padding_ must be a multiple of `lacunarity ** (octaves - 1) * res`. Most latent types will have one latent pixel equaling eight normal pixels - i.e. if your image is 512x512, the latent would be 64x64.
|
||||
|
||||
#### Node Parameters
|
||||
|
||||
* `depth`: When non-zero, 3D perlin noise will be generated.
|
||||
* `detail_level`: Controls the detail level of the noise when `break_pattern` is non-zero. No effect when using 100% raw Perlin noise.
|
||||
* `octaves`: Generally controls the detail level of the noise. Each octave involves generating a layer of noise so there is a performance cost to increasing octaves.
|
||||
* `persistence`: Controls how rough the generated noise is. Lower values will result in smoother noise, higher values will look more like Gaussian noise. Comma-separated list, multiple items will apply to octaves in sequence.
|
||||
* `lacunarity_height`: Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.
|
||||
* `lacunarity_width`: " "
|
||||
* `lacunarity_depth`: " "
|
||||
* `res_height`: Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.
|
||||
* `res_width`: " "
|
||||
* `res_depth`: " "
|
||||
* `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.
|
||||
* `initial_depth`: First zero-based depth index the noise generator will return. Only has an effect when depth is non-zero.
|
||||
* `wrap_depth`: If non-zero, instead of generating a new chunk of noise when the last slice is used will instead jump back to the specified zero-based depth index. Only has an effect when depth is non-zero. Since this is repeating the same noise, you may need to reduce `s_noise` in samplers especially if your `depth` value is low.
|
||||
* `max_depth`: Basically crops the depth dimension to the specified value (inclusive). Negative values start from the end, the default of -1 does no cropping. Only has an effect when depth is non-zero. The reason you might want to use this is changing `depth` will also effectively change the seed.
|
||||
* `tileable_height`: Makes the specified dimension tileable. (May or may not work correctly.)
|
||||
* `tileable_width`: " "
|
||||
* `tileable_depth`: " "
|
||||
* `blend`: Blending function used when generating Perlin noise. When set to values other than LERP may not work at all or may not actually generate Perlin noise. If you have `ComfyUI-bleh` there will be many more blending options (see [Integration](#integration)).
|
||||
* `pattern_break_blend`: Blending function used to blend pattern broken noise with raw noise. If you have `ComfyUI-bleh` there will be many more blending options (see [Integration](#integration)).
|
||||
* `depth_over_channels`: When disabled, each channel will have its own separate 3D noise pattern. When enabled, depth is multiplied by the number of channels and each channel is a slice of depth. Only has an effect when depth is non-zero.
|
||||
* `pad_height`: Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation.
|
||||
* `pad_width`: " "
|
||||
* `pad_depth`: " "
|
||||
* `initial_amplitude`: Controls the amplitude for the first octave. The amplitude gets multiplied by `persistence` after each octave.
|
||||
* `initial_frequency_height`: Controls the frequency for the first octave for the this axis. The frequency gets multiplied by `lacunarity` after each octave.
|
||||
* `initial_frequency_width`: " "
|
||||
* `initial_frequency_depth`: " "
|
||||
* `normalize`: Controls whether the output noise is normalized after generation.
|
||||
* `device`: Controls what device is used to generate the noise. GPU noise may be slightly faster but you will get different results on different GPUs.
|
||||
|
||||
+2
-6
@@ -1,5 +1,5 @@
|
||||
from .py import nodes
|
||||
from .py import custom_noise
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"OCS Sampler": nodes.SamplerNode,
|
||||
@@ -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"]
|
||||
|
||||
@@ -21,10 +21,6 @@ Symbols (simple string type) are defined using `'symbol_name` - note the solitar
|
||||
|
||||
`;` can be used to sequence operations. I.E. `exp1 ; exp2` evaluates `exp1`, then `exp2` and then result of the expression is whatever `exp2` returned.
|
||||
|
||||
`:=` is used to assign to a temporary variable (see `set_var` below).
|
||||
|
||||
The expression language supports a C/JavaScript style ternary operator: `condition ? true_branch : false_branch` is the equivalent of `if(condition, true_branch, false_branch)`.
|
||||
|
||||
Like Python, a parenthesized expression with a trailing comma can be used to create an empty tuple. Example: `(1,)`
|
||||
|
||||
## Filter Variables
|
||||
@@ -102,8 +98,6 @@ Available in model filters, with the exception of the `input` filter.
|
||||
| <td colspan=3 align=left>Boolean negation</td> |
|
||||
|⬤| `s_` | start:`I(null)`, end:`I(null)`, step:`I(null)` | `slice` |
|
||||
| <td colspan=3 align=left>Creates a slice object from the `start`, `end`, `step` values. See Numpy [s_](https://numpy.org/doc/stable/reference/generatednumpy.s_.html)</td> |
|
||||
|⬤| `set_var` | `SY`, `*` | `*` |
|
||||
| <td colspan=3 align=left>Sets a temporary variable to the specified value and returns the value. Alias for the `:=` assignment operator. <br/> **Example**: `test1 := 2; set_var('test2, 10); test1 * test2`</td> |
|
||||
|⬤| `unsafe_call` | `callable`, `*`\* | `*` |
|
||||
| <td colspan=3 align=left>Allows calling an arbitrary callable. <br/> **Example:** `unsafe_call(some_callable, arg1, arg2, kwarg1 :> 123)`</td>
|
||||
|
||||
@@ -135,32 +129,11 @@ Available in model filters, with the exception of the `input` filter.
|
||||
| <td colspan=3 align=left>Rolls a tensor along the specified dimensions. If amount is >= -1.0 and < 1.0 this will be interpreted as a percentage. <br/> **Example:** `t_roll(some_tensor, 10, (-2,))`</td> |
|
||||
|⬤| `t_scale` | tensor:`T`, scale:`SN \| NS`, mode:`SY(bicubic)`, absolute_scale:`B(false)` | `T` |
|
||||
| <td colspan=3 align=left>Scales a tensor. If scale is a tuple, it will be interpreted as `(height, width)`. When `absolute_scale` is not set, the scales will be interpreted as percentages otherwise absolute values will be used. <br/> Example: `t_scale(some_tensor, (0.75, 0.5), 'bilinear)`</td> |
|
||||
|⬤| `t_scale_nnlatentupscale` | tensor:`T`, mode:`SY(sd1)`, scale:`SN(2.0)` | `T` |
|
||||
| <td colspan=3 align=left>Available if you have [ComfyUi_NNLatentUpscale](https://github.com/Ttl/ComfyUi_NNLatentUpscale) installed. `mode` must be one of `sd1` or `sdxl`. `scale` should be between 1.0 and 2.0 (may or may not work out of that range).<br/> **Example:** `t_scale_nnlatentupscale(some_tensor, 'sdxl, 1.5)`</td> |
|
||||
|⬤| `t_shape` | tensor:`T` | `SN` |
|
||||
| <td colspan=3 align=left>Returns a tensor's shape as a tuple. <br/> Example: `shp := t_shape(some_tensor); width := shp[-1]; height := shp[-2]`</td> |
|
||||
|⬤| `t_std` | tensor:`T`, dim:`SN(-3, -2, -1)` | `T` |
|
||||
| <td colspan=3 align=left>Tensor std, second argument is dimensions. <br/> **Example:** `t_std(some_tensor, (-2, -1))`</td> |
|
||||
|⬤| `t_taesd_decode` | tensor:`T`, mode:`SY(sd15)` | `T` |
|
||||
| <td colspan=3 align=left>Decodes a latent tensor used TAESD. Mode must be one of `sd15`, `sdxl`. Only works if the appropriate models are in `vae_approx` <br/> **Example:** `t_taesd_decode(some_tensor, 'sd15)`</td> |
|
||||
|⬤| `unsafe_tensor_method` | `T`, `SY`, `*`\* | `*` |
|
||||
| <td colspan=3 align=left>Unsafe tensor method call. See note below. <br/> **Example:** `unsafe_tensor_method(some_tensor, 'mul, 10)`</td> |
|
||||
|⬤| `unsafe_torch` | path:`SY` | `*` |
|
||||
| <td colspan=3 align=left>Unsafe Torch module attribute access. See note below. <br/> **Example:** `unsafe_torch('nn.functional.interpolate)`</td> |
|
||||
|
||||
**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
|
||||
|
||||
`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.
|
||||
|
||||
| | Name | Input | Output |
|
||||
| :--- | :--- | :--- | :--- |
|
||||
|⬤| `img_pil_resize` | image:`IMG`, size:`SN \| NS`, resample_mode:`SY(bicubic)`, absolute_scale:`B(false)` | `IMG` |
|
||||
| <td colspan=3 align=left>Scales an image batch using Pillow's [`Image.resize`](https://pillow.readthedocs.io/en/stable/reference/Image.html#PIL.Image.Image.resize) function (follow link for information about resample modes, etc). If scale is a tuple, it will be interpreted as `(height, width)`. When `absolute_scale` is not set, the scales will be interpreted as percentages otherwise absolute values will be used. <br/> Example: `img_pil_resize(image_batch, (0.75, 0.5), 'lanczos)`</td> |
|
||||
|⬤| `img_shape` | image:`IMG` | `SN` |
|
||||
| <td colspan=3 align=left>Returns an image's shape as a tuple. Will fail if all the images in the batch aren't the same size. <br/> Example: `shp := img_shape(image_batch); width := shp[-1]; height := shp[-2]`</td> |
|
||||
|⬤| `img_taesd_encode` | image:`IMG`, reference_latent: `T`, mode:`SY(sd15)` | `T` |
|
||||
| <td colspan=3 align=left>Encodes an image batch into a latent tensor. The reference latent is only used to determine what device and type the output should be. Mode must be one of `sd15`, `sdxl`. Only works if the appropriate models are in `vae_approx`. <br/> **Example:** `img_taesd_encode(image_batch, some_tensor, 'sd15)`</td> |
|
||||
|
||||
@@ -92,8 +92,6 @@ final: default
|
||||
|
||||
There may be additional keys depending on the filter type.
|
||||
|
||||
If you have [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) available, you can use any blend mode it supports. Otherwise OCS provides these built-in blend modes: `lerp`, `a_only`, `b_only`. _Note_: `a` is considered the original value, `b` the changed value. `a_only` and `b_only` will still scale their output by the `strength`.
|
||||
|
||||
## Filter Types
|
||||
|
||||
### `simple`
|
||||
|
||||
@@ -1,11 +0,0 @@
|
||||
from . import nodes
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"OCSNoise PerlinSimple": nodes.PerlinSimpleNode,
|
||||
"OCSNoise PerlinAdvanced": nodes.PerlinAdvancedNode,
|
||||
"OCSNoise ImmiscibleReference": nodes.ImmiscibleReferenceNoiseNode,
|
||||
"OCSNoise to SONAR_CUSTOM_NOISE": nodes.ToSonarNode,
|
||||
"OCSNoise Conditioning": nodes.NoiseConditioningNode,
|
||||
"OCSNoise OverrideSamplerNoise": nodes.SamplerNodeConfigOverride,
|
||||
"OCSNoise ExpressionFilteredNoise": nodes.ExpressionFilteredNoiseNode,
|
||||
}
|
||||
@@ -1,180 +0,0 @@
|
||||
import abc
|
||||
from typing import Any, Callable
|
||||
|
||||
import torch
|
||||
|
||||
from ..external import IntegratedNode
|
||||
from ..nodes import NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE
|
||||
from ..noise import scale_noise
|
||||
|
||||
|
||||
class CustomNoiseItemBase(abc.ABC):
|
||||
def __init__(self, factor, **kwargs):
|
||||
self.factor = factor
|
||||
self.keys = set(kwargs.keys())
|
||||
for k, v in kwargs.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
def clone_key(self, k):
|
||||
return getattr(self, k)
|
||||
|
||||
def clone(self):
|
||||
return self.__class__(self.factor, **{k: self.clone_key(k) for k in self.keys})
|
||||
|
||||
def set_factor(self, factor):
|
||||
self.factor = factor
|
||||
return self
|
||||
|
||||
def get_normalize(self, k, default=None):
|
||||
val = getattr(self, k, None)
|
||||
return default if val is None else val
|
||||
|
||||
@abc.abstractmethod
|
||||
def make_noise_sampler(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
sigma_min=None,
|
||||
sigma_max=None,
|
||||
seed=None,
|
||||
cpu=True,
|
||||
normalized=True,
|
||||
):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class CustomNoiseChain:
|
||||
def __init__(self, items=None):
|
||||
self.items = items if items is not None else []
|
||||
|
||||
def clone(self):
|
||||
return CustomNoiseChain(
|
||||
[i.clone() for i in self.items],
|
||||
)
|
||||
|
||||
def add(self, item):
|
||||
if item is None:
|
||||
raise ValueError("Attempt to add nil item")
|
||||
self.items.append(item)
|
||||
|
||||
@property
|
||||
def factor(self):
|
||||
return sum(abs(i.factor) for i in self.items)
|
||||
|
||||
def rescaled(self, scale=1.0):
|
||||
divisor = self.factor / scale
|
||||
divisor = divisor if divisor != 0 else 1.0
|
||||
result = self.clone()
|
||||
if divisor != 1:
|
||||
for i in result.items:
|
||||
i.set_factor(i.factor / divisor)
|
||||
return result
|
||||
|
||||
@torch.no_grad()
|
||||
def make_noise_sampler(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
sigma_min=None,
|
||||
sigma_max=None,
|
||||
seed=None,
|
||||
cpu=True,
|
||||
normalized=True,
|
||||
) -> Callable:
|
||||
noise_samplers = tuple(
|
||||
i.make_noise_sampler(
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu,
|
||||
normalized=False,
|
||||
)
|
||||
for i in self.items
|
||||
)
|
||||
if not noise_samplers or not all(noise_samplers):
|
||||
raise ValueError("Failed to get noise sampler")
|
||||
factor = self.factor
|
||||
|
||||
def noise_sampler(sigma, sigma_next):
|
||||
result = None
|
||||
for ns in noise_samplers:
|
||||
noise = ns(sigma, sigma_next)
|
||||
if result is None:
|
||||
result = noise
|
||||
else:
|
||||
result += noise
|
||||
return scale_noise(result, factor, normalized=normalized)
|
||||
|
||||
return noise_sampler
|
||||
|
||||
|
||||
class CustomNoiseNodeBase(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "An Overly Complicated Sampling custom noise item."
|
||||
RETURN_TYPES = ("OCS_NOISE",)
|
||||
OUTPUT_TOOLTIPS = ("A custom noise chain.",)
|
||||
CATEGORY = "OveryComplicatedSampling/noise"
|
||||
FUNCTION = "go"
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_item_class(self):
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls, *, include_rescale=True, include_chain=True):
|
||||
result = {
|
||||
"required": {
|
||||
"factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": -100.0,
|
||||
"max": 100.0,
|
||||
"step": 0.001,
|
||||
"round": False,
|
||||
"tooltip": "Scaling factor for the generated noise of this type.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {},
|
||||
}
|
||||
if include_rescale:
|
||||
result["required"] |= {
|
||||
"rescale": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0,
|
||||
"step": 0.001,
|
||||
"round": False,
|
||||
"tooltip": "When non-zero, this custom noise item and other custom noise items items connected to it will have their factor scaled to add up to the specified rescale value.",
|
||||
},
|
||||
),
|
||||
}
|
||||
if include_chain:
|
||||
result["optional"] |= {
|
||||
"ocs_noise_opt": (
|
||||
WILDCARD_NOISE,
|
||||
{
|
||||
"tooltip": f"Optional input for more custom noise items.\n{NOISE_INPUT_TYPES_HINT}",
|
||||
},
|
||||
),
|
||||
}
|
||||
return result
|
||||
|
||||
def go(
|
||||
self,
|
||||
factor=1.0,
|
||||
rescale=0.0,
|
||||
ocs_noise_opt=None,
|
||||
**kwargs: dict[str, Any],
|
||||
):
|
||||
nis = ocs_noise_opt.clone() if ocs_noise_opt else CustomNoiseChain()
|
||||
if factor != 0:
|
||||
nis.add(self.get_item_class()(factor, **kwargs))
|
||||
return (nis if rescale == 0 else nis.rescaled(rescale),)
|
||||
|
||||
|
||||
class NormalizeNoiseNodeMixin:
|
||||
@staticmethod
|
||||
def get_normalize(val: str) -> None | bool:
|
||||
return None if val == "default" else val == "forced"
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -1,654 +0,0 @@
|
||||
# Initial revision based on Perlin generation routines from https://github.com/Extraltodeus/noise_latent_perlinpinpin which was based on https://gist.github.com/vadimkantorov/ac1b097753f217c5c11bc2ff396e0a57 which was based on https://github.com/pvigier/perlin-numpy
|
||||
import math
|
||||
from typing import Any, Callable, NamedTuple, Sequence
|
||||
|
||||
import torch
|
||||
from comfy import model_management
|
||||
from tqdm import tqdm
|
||||
|
||||
from .. import filtering
|
||||
from ..latent import normalize_to_scale
|
||||
from ..noise import scale_noise
|
||||
from .base import CustomNoiseItemBase, NormalizeNoiseNodeMixin
|
||||
|
||||
|
||||
def smoothstep_function(t):
|
||||
return 6 * t**5 - 15 * t**4 + 10 * t**3
|
||||
|
||||
|
||||
class BlendFunction(NamedTuple):
|
||||
name: str = "lerp"
|
||||
blend_function: Callable[..., torch.Tensor] = torch.lerp
|
||||
|
||||
def __call__(self, *args: Any, **kwargs: Any) -> torch.Tensor:
|
||||
return self.blend_function(*args, **kwargs)
|
||||
|
||||
|
||||
class Perlin(NamedTuple):
|
||||
depth: int = 16
|
||||
res: tuple[tuple[float, ...], ...] = ((1.0,), (1.0,), (1.0,))
|
||||
octaves: int = 2
|
||||
persistence: tuple[float, ...] = (1.0,)
|
||||
lacunarity: tuple[tuple[float, ...], ...] = ((2,), (2,), (2,))
|
||||
initial_amplitude: float = 1.0
|
||||
initial_frequency: tuple[float, ...] = (1.0, 1.0, 1.0)
|
||||
break_pattern: float = 0.99
|
||||
break_pattern_multiplier: float = 100000.0
|
||||
break_pattern_use_frac: bool = True
|
||||
detail_level: float = 0.0
|
||||
ridge_blend: BlendFunction = BlendFunction()
|
||||
ridge_weight: float = 0.0
|
||||
ridge_scale: float = 1.0
|
||||
warp_strength: float = 0.0
|
||||
octave_shift: float = 0.0
|
||||
curl_strength: float = 0.0
|
||||
curl_dims: tuple[int, int] = (0, 1)
|
||||
tileable: tuple[bool, ...] = (False, False, False)
|
||||
fade: Callable[[torch.Tensor], torch.Tensor] = smoothstep_function
|
||||
blend: BlendFunction = BlendFunction()
|
||||
pattern_break_blend: BlendFunction = BlendFunction()
|
||||
depth_over_channels: bool = False
|
||||
initial_depth: int = 0
|
||||
wrap_depth: int = 0
|
||||
max_depth: int = -1
|
||||
pad: tuple[int, ...] = (0, 0, 0)
|
||||
pad_mode: str = "replicate"
|
||||
generator: torch.Generator | None = None
|
||||
device: str | torch.device = "default"
|
||||
dtype: torch.dtype | None = None
|
||||
|
||||
@classmethod
|
||||
def build(cls, **kwargs: Any) -> "Perlin":
|
||||
dfl = cls()
|
||||
depth = kwargs.get("depth", dfl.depth)
|
||||
for bk in ("blend", "pattern_break_blend", "ridge_blend"):
|
||||
bv = kwargs.pop(bk, getattr(dfl, bk))
|
||||
kwargs[bk] = (
|
||||
BlendFunction(bv, filtering.BLENDING_MODES[bv])
|
||||
if isinstance(bv, str)
|
||||
else bv
|
||||
)
|
||||
lacunarity = kwargs.pop("lacunarity", None)
|
||||
kwargs["lacunarity"] = (
|
||||
cls.maybe_parse_dhw_triple(
|
||||
(
|
||||
kwargs.pop("lacunarity_depth", dfl.lacunarity[0]),
|
||||
kwargs.pop("lacunarity_height", dfl.lacunarity[1]),
|
||||
kwargs.pop("lacunarity_width", dfl.lacunarity[2]),
|
||||
),
|
||||
depth,
|
||||
)
|
||||
if lacunarity is None
|
||||
else tuple(lacunarity)
|
||||
)
|
||||
res = kwargs.pop("res", None)
|
||||
kwargs["res"] = (
|
||||
cls.maybe_parse_dhw_triple(
|
||||
(
|
||||
kwargs.pop("res_depth", dfl.res[0]),
|
||||
kwargs.pop("res_height", dfl.res[1]),
|
||||
kwargs.pop("res_width", dfl.res[2]),
|
||||
),
|
||||
depth,
|
||||
)
|
||||
if res is None
|
||||
else tuple(res)
|
||||
)
|
||||
pad = kwargs.pop("pad", None)
|
||||
kwargs["pad"] = (
|
||||
(
|
||||
kwargs.pop("pad_depth", dfl.pad[0]),
|
||||
kwargs.pop("pad_height", dfl.pad[1]),
|
||||
kwargs.pop("pad_width", dfl.pad[2]),
|
||||
)
|
||||
if pad is None
|
||||
else tuple(pad)
|
||||
)
|
||||
initial_frequency = kwargs.pop("initial_frequency", None)
|
||||
kwargs["initial_frequency"] = (
|
||||
(
|
||||
kwargs.pop("initial_frequency_depth", dfl.initial_frequency[0]),
|
||||
kwargs.pop("initial_frequency_height", dfl.initial_frequency[1]),
|
||||
kwargs.pop("initial_frequency_width", dfl.initial_frequency[2]),
|
||||
)[int(depth == 0) :]
|
||||
if initial_frequency is None
|
||||
else tuple(initial_frequency)
|
||||
)
|
||||
tileable = kwargs.pop("tileable", None)
|
||||
kwargs["tileable"] = (
|
||||
(
|
||||
kwargs.pop("tileable_depth", dfl.tileable[0]),
|
||||
kwargs.pop("tileable_height", dfl.tileable[1]),
|
||||
kwargs.pop("tileable_width", dfl.tileable[2]),
|
||||
)[int(depth == 0) :]
|
||||
if tileable is None
|
||||
else tuple(tileable)
|
||||
)
|
||||
persistence = kwargs.pop("persistence", None)
|
||||
if persistence is not None:
|
||||
kwargs["persistence"] = (
|
||||
cls.maybe_parse_commasep_list(persistence)
|
||||
if isinstance(persistence, str)
|
||||
else tuple(persistence)
|
||||
)
|
||||
curl_dims = kwargs.pop("curl_dims", None)
|
||||
if curl_dims is not None:
|
||||
kwargs["curl_dims"] = tuple(
|
||||
int(v)
|
||||
for v in (
|
||||
cls.maybe_parse_commasep_list(curl_dims)
|
||||
if isinstance(curl_dims, str)
|
||||
else curl_dims
|
||||
)
|
||||
)
|
||||
fs = frozenset(cls._fields)
|
||||
kwargs = {k: v for k, v in kwargs.items() if k in fs}
|
||||
return cls(**kwargs)
|
||||
|
||||
@classmethod
|
||||
def maybe_parse_dhw_triple(cls, val, depth, convert=float):
|
||||
return tuple(cls.maybe_parse_commasep_list(v) for v in val)[int(depth == 0) :]
|
||||
|
||||
@classmethod
|
||||
def maybe_parse_commasep_list(cls, val, convert=float):
|
||||
if not isinstance(val, str):
|
||||
return val
|
||||
return tuple(convert(v) for v in val.strip().split(",") if v.strip())
|
||||
|
||||
def get_commasep(self, key, idx=None):
|
||||
val = getattr(self, key)
|
||||
if idx is not None:
|
||||
val = val[idx]
|
||||
return ", ".join(repr(v) for v in val)
|
||||
|
||||
def octave(
|
||||
self,
|
||||
shape: Sequence[int],
|
||||
res: Sequence[float],
|
||||
*,
|
||||
batch_size: int = 1,
|
||||
channels: int = 1,
|
||||
octave: int = 0,
|
||||
base_noise: torch.Tensor | None = None,
|
||||
warp: torch.Tensor | None = None,
|
||||
):
|
||||
shape = tuple(shape)
|
||||
res = tuple(res)
|
||||
dims = len(res)
|
||||
didxs = tuple(range(dims))
|
||||
|
||||
coords = tuple(
|
||||
torch.linspace(
|
||||
0,
|
||||
res[i],
|
||||
shape[i] + 1,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)[:-1]
|
||||
for i in didxs
|
||||
)
|
||||
grid_coords = torch.meshgrid(*coords, indexing="ij")
|
||||
p = torch.stack(grid_coords, dim=-1)
|
||||
|
||||
# Expand `p` to include Batch and Channel dimensions
|
||||
p = (
|
||||
p.unsqueeze(0)
|
||||
.unsqueeze(0)
|
||||
.expand(
|
||||
batch_size,
|
||||
channels,
|
||||
*((-1,) * (dims + 1)),
|
||||
)
|
||||
)
|
||||
|
||||
# Domain Warping: Apply warp before calculating p0 and grid.
|
||||
if warp is not None and self.warp_strength != 0.0:
|
||||
p += warp.unsqueeze(-1) * self.warp_strength
|
||||
|
||||
# Now calculate indices and bounds safely
|
||||
p0 = p.floor().long()
|
||||
grid = p - p0
|
||||
grad_shape = tuple(int(math.ceil(res[i])) + 1 for i in didxs)
|
||||
|
||||
gradients = torch.randn(
|
||||
batch_size,
|
||||
channels,
|
||||
*grad_shape,
|
||||
dims,
|
||||
generator=self.generator,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
gradients = torch.nn.functional.normalize(gradients, dim=-1)
|
||||
|
||||
# Modulate Perlin amplitude using base noise
|
||||
if base_noise is not None:
|
||||
gradients = gradients * base_noise.unsqueeze(-1).to(gradients)
|
||||
|
||||
octave_shift = round((1.0 + octave) * self.octave_shift)
|
||||
if octave_shift != 0:
|
||||
gradients = gradients.roll(dims=-1, shifts=octave_shift)
|
||||
|
||||
if dims > 1 and self.curl_strength != 0:
|
||||
# Generate a random spin angle for every single gradient point on the grid
|
||||
angles = torch.randn(
|
||||
batch_size,
|
||||
channels,
|
||||
*grad_shape,
|
||||
generator=self.generator,
|
||||
device=gradients.device,
|
||||
dtype=gradients.dtype,
|
||||
).mul_(self.curl_strength)
|
||||
|
||||
cos_a = angles.cos()
|
||||
sin_a = angles.sin_()
|
||||
|
||||
# Grab the first two axes by default (e.g., Depth/Height, or Height/Width)
|
||||
d1, d2 = self.curl_dims[:2]
|
||||
g0 = gradients[..., d1].clone()
|
||||
g1 = gradients[..., d2].clone()
|
||||
|
||||
# Apply 2D Rotation Matrix to twist the vectors
|
||||
gradients[..., d1] = (g0 * cos_a).sub_(g1 * sin_a)
|
||||
gradients[..., d2] = (g0 * sin_a).add_(g1 * cos_a)
|
||||
|
||||
def get_shift(n, dims, *, on_value, off_value):
|
||||
return tuple(
|
||||
on_value if n & (1 << bitidx) else off_value for bitidx in range(dims)
|
||||
)
|
||||
|
||||
def blend_reduce(vals, t, depth=0):
|
||||
curr_t = t[..., depth]
|
||||
if len(vals) == 2:
|
||||
return self.blend(*vals, curr_t)
|
||||
pairs = zip(vals[0::2], vals[1::2])
|
||||
return blend_reduce(
|
||||
tuple(self.blend(v1, v2, curr_t) for v1, v2 in pairs), t, depth + 1
|
||||
)
|
||||
|
||||
ns = []
|
||||
|
||||
# Pre-calculate batched indices for tensor indexing
|
||||
b_idx = torch.arange(batch_size, device=self.device).view(
|
||||
batch_size, 1, *[1] * dims
|
||||
)
|
||||
c_idx = torch.arange(channels, device=self.device).view(
|
||||
1, channels, *[1] * dims
|
||||
)
|
||||
|
||||
for i in range(1 << dims):
|
||||
shift = get_shift(i, dims, off_value=0, on_value=1)
|
||||
|
||||
idx = p0.clone()
|
||||
for dim in range(dims):
|
||||
idx[..., dim] += shift[dim]
|
||||
idx[..., dim] %= grad_shape[dim] - int(self.tileable[dim])
|
||||
|
||||
spatial_indices = tuple(idx[..., dim] for dim in range(dims))
|
||||
grad = gradients[(b_idx, c_idx) + spatial_indices]
|
||||
|
||||
grid_shift = get_shift(i, dims, off_value=0, on_value=-1)
|
||||
grid_shift_tensor = torch.tensor(
|
||||
grid_shift, dtype=self.dtype, device=self.device
|
||||
)
|
||||
|
||||
d = ((grid + grid_shift_tensor) * grad).sum(dim=-1)
|
||||
ns.append(d)
|
||||
|
||||
return blend_reduce(ns, self.fade(grid)).mul_(2.0**0.5)
|
||||
|
||||
@staticmethod
|
||||
def get_wrap_dim(val, *dims):
|
||||
for dim in dims:
|
||||
nelem = len(val) if not isinstance(val, torch.Tensor) else val.shape[0]
|
||||
val = val[dim % nelem]
|
||||
return val
|
||||
|
||||
def get_unwrapped_octaves_dims(self, val, ndim: int) -> torch.Tensor:
|
||||
return torch.tensor(
|
||||
tuple(
|
||||
self.get_wrap_dim(val, didx, oidx)
|
||||
for oidx in range(self.octaves)
|
||||
for didx in range(ndim)
|
||||
),
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
).reshape(self.octaves, ndim)
|
||||
|
||||
def generate_octaves(
|
||||
self,
|
||||
shape: Sequence[int],
|
||||
*,
|
||||
batch_size: int = 1,
|
||||
channels: int = 1,
|
||||
base_noise: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
shape = tuple(shape)
|
||||
ndim = len(shape)
|
||||
|
||||
amplitude = self.initial_amplitude
|
||||
res = self.get_unwrapped_octaves_dims(self.res, ndim)
|
||||
lacunarity = self.get_unwrapped_octaves_dims(self.lacunarity, ndim)
|
||||
persistence = self.persistence[: self.octaves]
|
||||
previous_octave = None
|
||||
|
||||
initial_frequency = self.initial_frequency[-ndim:]
|
||||
frequency = torch.ones(ndim, dtype=self.dtype, device=self.device)
|
||||
frequency[: len(initial_frequency)] = frequency.new(initial_frequency)
|
||||
noise = torch.zeros(
|
||||
(batch_size, channels, *shape),
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
for octave in range(self.octaves):
|
||||
octave_res = tuple(
|
||||
frequency[didx].item() * res[octave][didx].item()
|
||||
for didx in range(ndim)
|
||||
)
|
||||
|
||||
grad_shape = tuple(int(math.ceil(octave_res[i])) + 1 for i in range(ndim))
|
||||
|
||||
if base_noise is None:
|
||||
octave_base_noise = None
|
||||
else:
|
||||
if base_noise.shape[-len(grad_shape) :] != grad_shape:
|
||||
mode = (
|
||||
"bilinear"
|
||||
if ndim == 2
|
||||
else ("trilinear" if ndim == 3 else "nearest")
|
||||
)
|
||||
octave_base_noise = torch.nn.functional.interpolate(
|
||||
base_noise,
|
||||
size=grad_shape,
|
||||
mode=mode,
|
||||
**(
|
||||
{"align_corners": False}
|
||||
if mode not in ("nearest", "area")
|
||||
else {}
|
||||
),
|
||||
)
|
||||
else:
|
||||
octave_base_noise = base_noise
|
||||
|
||||
octave_output = self.octave(
|
||||
shape,
|
||||
octave_res,
|
||||
batch_size=batch_size,
|
||||
channels=channels,
|
||||
octave=octave,
|
||||
base_noise=octave_base_noise,
|
||||
warp=previous_octave,
|
||||
)
|
||||
|
||||
if self.ridge_weight != 0.0:
|
||||
ridge = (
|
||||
1.0
|
||||
- octave_output.div(
|
||||
octave_output.abs().max().clamp_min_(1e-07)
|
||||
).abs_()
|
||||
)
|
||||
ridge -= 0.5
|
||||
ridge *= 2.0 * self.ridge_scale
|
||||
octave_output = self.ridge_blend(
|
||||
octave_output,
|
||||
ridge,
|
||||
self.ridge_weight,
|
||||
)
|
||||
|
||||
noise += amplitude * octave_output
|
||||
previous_octave = octave_output
|
||||
|
||||
frequency *= lacunarity[octave]
|
||||
amplitude *= self.get_wrap_dim(persistence, octave)
|
||||
|
||||
return noise
|
||||
|
||||
# Based on approach from https://github.com/Extraltodeus/noise_latent_perlinpinpin
|
||||
@staticmethod
|
||||
def break_pattern_func(
|
||||
t: torch.Tensor,
|
||||
detail: float = 0.0,
|
||||
*,
|
||||
multiplier: float = 1000000.0,
|
||||
use_frac: bool = False,
|
||||
clamp_low: float = -5.0,
|
||||
clamp_high: float = 5.0,
|
||||
) -> torch.Tensor:
|
||||
detail_factor = (1 + detail * 0.1) * 2.0**0.5 * 0.2
|
||||
result = t.abs().mul_(multiplier)
|
||||
result = result.frac_() if use_frac else result.remainer_(11).div_(11)
|
||||
return (
|
||||
result.mul_(2)
|
||||
.sub_(1)
|
||||
.erfinv_()
|
||||
.mul_(detail_factor)
|
||||
.clamp_(clamp_low, clamp_high)
|
||||
)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
width: int,
|
||||
height: int,
|
||||
*,
|
||||
batch_size: int = 1,
|
||||
channels: int = 4,
|
||||
base_noise: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
depth = self.depth
|
||||
pad_depth, pad_height, pad_width = self.pad[:3]
|
||||
depth_over_channels = self.depth_over_channels
|
||||
|
||||
if depth < 1:
|
||||
depth_over_channels = False
|
||||
pad_depth = 0
|
||||
eff_shape = (height + pad_height * 2, width + pad_width * 2)
|
||||
eff_channels = channels
|
||||
eff_depth = 0
|
||||
noise_dims = 2
|
||||
else:
|
||||
eff_channels = channels if not depth_over_channels else 1
|
||||
eff_depth = depth if not depth_over_channels else depth * channels
|
||||
eff_shape = (
|
||||
eff_depth + pad_depth * 2,
|
||||
height + pad_height * 2,
|
||||
width + pad_width * 2,
|
||||
)
|
||||
noise_dims = 3
|
||||
|
||||
bn = base_noise
|
||||
if bn is not None:
|
||||
if depth_over_channels and depth > 0:
|
||||
# Flattens 5D base_noise into a contiguous sequential depth format matching outputs
|
||||
# (B, C, D, H, W) -> (B, 1, D*C, H, W)
|
||||
bn = bn.movedim(1, 2).reshape(batch_size, 1, eff_depth, height, width)
|
||||
|
||||
if pad_width > 0 or pad_height > 0 or pad_depth > 0:
|
||||
if noise_dims == 3:
|
||||
pad_tuple = (
|
||||
pad_width,
|
||||
pad_width,
|
||||
pad_height,
|
||||
pad_height,
|
||||
pad_depth,
|
||||
pad_depth,
|
||||
)
|
||||
else:
|
||||
pad_tuple = (pad_width, pad_width, pad_height, pad_height)
|
||||
bn = torch.nn.functional.pad(bn, pad_tuple, mode=self.pad_mode)
|
||||
|
||||
noise_values = self.generate_octaves(
|
||||
eff_shape,
|
||||
batch_size=batch_size,
|
||||
channels=eff_channels,
|
||||
base_noise=bn,
|
||||
)
|
||||
|
||||
# Apply normalization to the spatial dimensions individually per batch and channel
|
||||
norm_dims = tuple(range(-len(eff_shape), 0))
|
||||
noise_values = normalize_to_scale(noise_values, -1.0, 1.0, dim=norm_dims)
|
||||
|
||||
if self.break_pattern != 0.0:
|
||||
result = self.pattern_break_blend(
|
||||
noise_values,
|
||||
self.break_pattern_func(
|
||||
noise_values,
|
||||
detail=self.detail_level,
|
||||
use_frac=self.break_pattern_use_frac,
|
||||
multiplier=self.break_pattern_multiplier,
|
||||
),
|
||||
self.break_pattern,
|
||||
)
|
||||
else:
|
||||
result = noise_values
|
||||
|
||||
if sum(self.pad[:3]) > 0:
|
||||
if noise_dims == 3:
|
||||
result = result[
|
||||
:,
|
||||
:,
|
||||
pad_depth : eff_depth + pad_depth,
|
||||
pad_height : height + pad_height,
|
||||
pad_width : width + pad_width,
|
||||
]
|
||||
else:
|
||||
result = result[
|
||||
:,
|
||||
:,
|
||||
pad_height : height + pad_height,
|
||||
pad_width : width + pad_width,
|
||||
]
|
||||
|
||||
if depth_over_channels and depth > 0:
|
||||
# Map the contiguous sequential flattened depth back identically matching (D0_C0, D0_C1, ...) interleaving
|
||||
# (B, 1, D*C, H, W) -> (B, C, D, H, W)
|
||||
result = result.reshape(
|
||||
batch_size,
|
||||
depth,
|
||||
channels,
|
||||
height,
|
||||
width,
|
||||
).movedim(1, 2)
|
||||
|
||||
if noise_dims == 3:
|
||||
# Shift Depth index backwards ahead of Batch (per expectation of the make_noise_sampler)
|
||||
# (B, C, D, H, W) -> (D, B, C, H, W)
|
||||
result = result.movedim(2, 0)
|
||||
|
||||
return result.contiguous()
|
||||
|
||||
|
||||
class PerlinItem(CustomNoiseItemBase):
|
||||
def __init__(
|
||||
self,
|
||||
factor,
|
||||
*,
|
||||
perlin: Perlin | None = None,
|
||||
device=None,
|
||||
normalized=None,
|
||||
base_noise_opt=None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
if perlin is None:
|
||||
perlin = Perlin.build(**kwargs)
|
||||
super().__init__(
|
||||
factor,
|
||||
perlin=perlin,
|
||||
device=device,
|
||||
normalized=normalized
|
||||
if not isinstance(normalized, str)
|
||||
else NormalizeNoiseNodeMixin.get_normalize(normalized),
|
||||
base_noise_opt=base_noise_opt.clone()
|
||||
if base_noise_opt is not None
|
||||
else None,
|
||||
# **kwargs,
|
||||
)
|
||||
|
||||
def make_noise_sampler(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
sigma_min: float | None,
|
||||
sigma_max: float | None,
|
||||
seed: int | None,
|
||||
cpu: bool = True,
|
||||
normalized=True,
|
||||
) -> torch.Tensor: # ty:ignore[invalid-method-override]
|
||||
normalized = self.get_normalize("normalized", normalized)
|
||||
cpu = cpu if self.device == "default" and cpu else self.device == "cpu"
|
||||
device = torch.device("cpu") if cpu else model_management.get_torch_device()
|
||||
perlin: Perlin = self.perlin._replace(device=device, dtype=x.dtype)
|
||||
noise_chunk = None
|
||||
noise_index = perlin.initial_depth
|
||||
max_idx = None
|
||||
if x.ndim < 4:
|
||||
raise ValueError("Can only handle latents with 4+ dimensions")
|
||||
orig_shape = x.shape
|
||||
b = orig_shape[0]
|
||||
c = math.prod(orig_shape[1:-2]) # Hack to deal with video models
|
||||
h, w = orig_shape[-2:]
|
||||
|
||||
base_noise_sampler = (
|
||||
self.base_noise_opt.make_noise_sampler(
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu,
|
||||
normalized=False,
|
||||
)
|
||||
if self.base_noise_opt
|
||||
else None
|
||||
)
|
||||
|
||||
x_device, x_dtype = x.device, x.dtype
|
||||
del x
|
||||
|
||||
def noise_sampler(s, sn):
|
||||
nonlocal noise_chunk, noise_index, max_idx
|
||||
if noise_chunk is None:
|
||||
base_noise = None
|
||||
if base_noise_sampler:
|
||||
bn_tuple = tuple(
|
||||
base_noise_sampler(s, sn).reshape(b, c, h, w)
|
||||
for _ in range(max(1, perlin.depth))
|
||||
)
|
||||
base_noise = (
|
||||
bn_tuple[0]
|
||||
if perlin.depth < 1
|
||||
else torch.stack(bn_tuple, dim=0).movedim(0, 2)
|
||||
)
|
||||
del bn_tuple
|
||||
|
||||
noise_chunk = perlin(
|
||||
w,
|
||||
h,
|
||||
batch_size=b,
|
||||
channels=c,
|
||||
base_noise=base_noise,
|
||||
).to(device=x_device, dtype=x_dtype)
|
||||
|
||||
if perlin.depth < 1:
|
||||
noise = noise_chunk
|
||||
noise_chunk = None
|
||||
return scale_noise(noise, self.factor, normalized=normalized)
|
||||
if perlin.max_depth != 0 and perlin.max_depth != -1:
|
||||
noise_chunk = noise_chunk[: perlin.max_depth]
|
||||
chunk_shape = noise_chunk.shape
|
||||
max_idx = (
|
||||
chunk_shape[0] - 1
|
||||
if perlin.wrap_depth == 0
|
||||
else min(perlin.wrap_depth, chunk_shape[0] - 1)
|
||||
)
|
||||
if max_idx < 0:
|
||||
max_idx += chunk_shape[0]
|
||||
|
||||
noise = noise_chunk[noise_index]
|
||||
noise_index += 1
|
||||
if noise_index > max_idx:
|
||||
noise_index = 0
|
||||
if not perlin.wrap_depth:
|
||||
noise_chunk = None
|
||||
result = scale_noise(noise, self.factor, normalized=normalized)
|
||||
return result.reshape(orig_shape) if result.shape != orig_shape else result
|
||||
|
||||
return noise_sampler
|
||||
+11
-16
@@ -1,23 +1,18 @@
|
||||
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
|
||||
|
||||
__all__ = (
|
||||
"types",
|
||||
"expression",
|
||||
"handler",
|
||||
"util",
|
||||
"validation",
|
||||
"ValidateArg",
|
||||
"Expression",
|
||||
"Arg",
|
||||
"BaseHandler",
|
||||
"BASIC_HANDLERS",
|
||||
"expression",
|
||||
"Expression",
|
||||
"handler",
|
||||
"HandlerContext",
|
||||
"types",
|
||||
"util",
|
||||
"ValidateArg",
|
||||
"validation",
|
||||
)
|
||||
|
||||
+34
-82
@@ -1,33 +1,27 @@
|
||||
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
|
||||
|
||||
|
||||
class Expression:
|
||||
EXPR_RE = re.compile(
|
||||
r"""
|
||||
\s*
|
||||
(
|
||||
\d+ # Numeric literal
|
||||
\d+ # Possibly negative numeric literal
|
||||
(?: \. \d* )? # Floating point
|
||||
(?: e [+-] \d+)? # Scientific notation
|
||||
| (?: \*\* | // ) # Doubled operators
|
||||
@@ -36,13 +30,10 @@ class Expression:
|
||||
| (?: \|\| | && ) # Logic
|
||||
| [-+*/|!(),] # Operators
|
||||
| :> # Key value binop
|
||||
| := # Assignment
|
||||
| ; # Sequencing
|
||||
| :: # Method call
|
||||
| [?:] # Ternary
|
||||
| \[ | ] # Index
|
||||
| \.\.\. # Index ellipsis
|
||||
| '[-\w.:=]+ # Symbol
|
||||
| ;
|
||||
| \[ | ]
|
||||
| \.\.\.
|
||||
| '[\w.]+ # Symbol
|
||||
| `?[a-z][\w.]*`? # Function/variable names
|
||||
)
|
||||
\s*
|
||||
@@ -62,14 +53,10 @@ class Expression:
|
||||
return self.eval(*args, **kwargs)
|
||||
|
||||
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,29 +76,24 @@ 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))
|
||||
|
||||
|
||||
CONST_OP_HANDLERS = {
|
||||
STATIC_OP_HANDLERS = {
|
||||
"+": operator.add,
|
||||
"-": operator.sub,
|
||||
"*": operator.mul,
|
||||
@@ -126,29 +108,18 @@ CONST_OP_HANDLERS = {
|
||||
"idiv": operator.floordiv,
|
||||
"pow": operator.pow,
|
||||
"mod": operator.mod,
|
||||
"neg": operator.neg,
|
||||
">": operator.gt,
|
||||
"<": operator.lt,
|
||||
">=": operator.ge,
|
||||
"<=": operator.le,
|
||||
"!=": operator.ne,
|
||||
"==": operator.eq,
|
||||
}
|
||||
|
||||
|
||||
def is_const_value(val):
|
||||
return val in (None, True, False) or isinstance(val, (int, float, ExpSym))
|
||||
|
||||
|
||||
def make_funap(op, args=(), kwargs=None):
|
||||
if kwargs is None:
|
||||
kwargs = ExpDict()
|
||||
argc = len(args)
|
||||
if argc > 2 or len(kwargs) or not all(is_const_value(v) for v in args):
|
||||
if argc > 2 or len(kwargs) or not all(isinstance(v, (int, float)) for v in args):
|
||||
return ExpFunAp(op, args, kwargs)
|
||||
if argc == 1 and op in "-+":
|
||||
return -args[0] if op == "-" else args[0]
|
||||
h = CONST_OP_HANDLERS.get(op)
|
||||
h = STATIC_OP_HANDLERS.get(op)
|
||||
if h is None:
|
||||
return ExpFunAp(op, args, kwargs)
|
||||
return h(*args)
|
||||
@@ -163,9 +134,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):
|
||||
@@ -200,20 +171,12 @@ class ExprParserSpec(ParserSpec):
|
||||
raise ParseError(f"{left!r} is not a valid function/variable name")
|
||||
args = []
|
||||
while p.lexer and p.token != ")":
|
||||
args.append(p.parse_until(COMMA_PRECEDENCE))
|
||||
args.append(p.parse_until(1))
|
||||
if p.token == ",":
|
||||
p.advance()
|
||||
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 == ")":
|
||||
@@ -223,9 +186,15 @@ class ExprParserSpec(ParserSpec):
|
||||
|
||||
@staticmethod
|
||||
def left_semicolon(p, token, left, bp):
|
||||
r = None if p.token in (None, ")", ";") else p.parse_until(0)
|
||||
if p.token == ")" or p.token is None:
|
||||
return (
|
||||
left
|
||||
if isinstance(left, ExpStatements)
|
||||
else ExpStatements(ExpTuple((left,)))
|
||||
)
|
||||
r = p.parse_until(bp)
|
||||
return ExpStatements(
|
||||
ExpTuple((*left.statements, r))
|
||||
ExpTuple(*left.statements, r)
|
||||
if isinstance(left, ExpStatements)
|
||||
else ExpTuple((left, r))
|
||||
)
|
||||
@@ -236,20 +205,6 @@ class ExprParserSpec(ParserSpec):
|
||||
p.expect("]")
|
||||
return make_funap("index", ExpTuple((idx, left)))
|
||||
|
||||
@staticmethod
|
||||
def left_assign(p, token, left, bp):
|
||||
if not isinstance(left, (ExpOp, ExpSym)):
|
||||
raise ParseError(f"bad LHS type for assignment operation {type(left)}")
|
||||
val = p.parse_until(bp)
|
||||
return make_funap("set_var", ExpTuple((ExpSym(left), val)))
|
||||
|
||||
@staticmethod
|
||||
def left_ternary(p, token, left, bp):
|
||||
true_branch = p.parse_until(0)
|
||||
p.expect(":")
|
||||
false_branch = p.parse_until(bp)
|
||||
return make_funap("if", ExpTuple((left, true_branch, false_branch)))
|
||||
|
||||
@staticmethod
|
||||
def get_type(token):
|
||||
if isinstance(token, (int, float)):
|
||||
@@ -265,7 +220,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, ("*", "/"))
|
||||
@@ -276,12 +230,10 @@ class ExprParserSpec(ParserSpec):
|
||||
self.add_left(9, self.left_binop, ("&&",))
|
||||
self.add_left(7, self.left_binop, ("||",))
|
||||
self.add_left(6, self.left_kv, (":>",))
|
||||
self.add_leftright(5, self.left_ternary, ("?",))
|
||||
self.add_leftright(4, self.left_assign, (":=",))
|
||||
self.add_left(COMMA_PRECEDENCE, self.left_comma, (",",))
|
||||
self.add_left(1, self.left_semicolon, (";",))
|
||||
self.add_left(5, self.left_semicolon, (";",))
|
||||
self.add_left(1, self.left_comma, (",",))
|
||||
self.add_null(0, self.null_paren, ("(",))
|
||||
self.add_null(
|
||||
-1, self.null_constant, ("number", "op", "sym", Ellipsis, True, False, None)
|
||||
)
|
||||
self.add_null(-1, ParserSpec.null_error, (")", "]", ":"))
|
||||
self.add_null(-1, ParserSpec.null_error, (")", "]"))
|
||||
|
||||
+25
-176
@@ -1,58 +1,14 @@
|
||||
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
|
||||
from .util import torch
|
||||
from .validation import Arg, ValidateArg, ValidateError
|
||||
|
||||
|
||||
class HandlerError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class HandlerContext:
|
||||
def __init__(self, handlers=None, constants=None, variables=None):
|
||||
self.handlers = handlers if handlers is not None else {}
|
||||
self.constants = constants if constants is not None else {}
|
||||
self.variables = variables if variables is not None else {}
|
||||
|
||||
def get_handler(self, k, default=Empty):
|
||||
return self.handlers.get(k, default)
|
||||
|
||||
def get_var(self, k, default=Empty):
|
||||
result = self.constants.get(k, Empty)
|
||||
if result is Empty:
|
||||
result = self.variables.get(k, Empty)
|
||||
return default if result is Empty else result
|
||||
|
||||
def set_var(self, k, v):
|
||||
if k in self.constants:
|
||||
raise KeyError(
|
||||
f"Cannot set variable with key {k}: already exists as a constant"
|
||||
)
|
||||
self.variables[k] = v
|
||||
|
||||
def unset_var(self, k):
|
||||
if k in self.variables:
|
||||
del self.variables[k]
|
||||
return True
|
||||
return False
|
||||
|
||||
def __contains__(self, k):
|
||||
return any(
|
||||
k in coll for coll in (self.handlers, self.constants, self.variables)
|
||||
)
|
||||
|
||||
def clone(self, *, handlers=Empty, constants=Empty, variables=Empty):
|
||||
return self.__class__(
|
||||
self.handlers if handlers is Empty else handlers,
|
||||
self.constants if constants is Empty else constants,
|
||||
self.variables if variables is Empty else variables,
|
||||
)
|
||||
|
||||
|
||||
class BaseHandler:
|
||||
input_validators = ()
|
||||
|
||||
@@ -64,12 +20,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)
|
||||
@@ -97,7 +50,7 @@ class BaseHandler:
|
||||
str_eff_key = True
|
||||
else:
|
||||
raise ValidateError(
|
||||
f"Error validating input argument {key} for {obj.name}, out of range for actual function arguments"
|
||||
f"Error validating input argument {key}, out of range for actual function arguments"
|
||||
)
|
||||
if getter is None:
|
||||
if str_eff_key:
|
||||
@@ -112,8 +65,8 @@ class BaseHandler:
|
||||
return validator(key, val)
|
||||
except ValidateError as exc:
|
||||
raise ValidateError(
|
||||
f"Error validating input argument {key} for {obj.name}, type {type(val)}: {exc!r}"
|
||||
)
|
||||
f"Error validating input argument {key}, type {type(val)}: {exc!r}"
|
||||
) from None
|
||||
|
||||
def safe_get_multi(self, keys, obj, getter=None, *, default=Empty):
|
||||
return (self.safe_get(k, obj, getter, default=default) for k in keys)
|
||||
@@ -184,7 +137,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 +171,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,28 +213,12 @@ 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"),)
|
||||
|
||||
def handle(self, obj, getter):
|
||||
key = self.safe_get(0, obj, getter=getter)
|
||||
return key in getter.ctx
|
||||
return key in getter.handlers
|
||||
|
||||
def validate_output(self, obj, value):
|
||||
return operator.truth(value)
|
||||
@@ -297,25 +232,17 @@ class GetHandler(BaseHandler):
|
||||
|
||||
def handle(self, obj, getter):
|
||||
key = self.safe_get("name", obj, getter=getter)
|
||||
result = getter.ctx.get_var(key)
|
||||
if result is Empty:
|
||||
h = getter.handlers.get(key)
|
||||
if h is None:
|
||||
return self.safe_get("fallback", obj, getter=getter)
|
||||
return ExpOp(key).eval(getter.ctx, *getter.args, **getter.kwargs)
|
||||
return h(getter.handlers, *getter.args, **getter.kwargs)
|
||||
|
||||
|
||||
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 +277,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 +287,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)
|
||||
@@ -385,62 +305,6 @@ class CommentHandler(BaseHandler):
|
||||
return None
|
||||
|
||||
|
||||
class SetVarHandler(BaseHandler):
|
||||
input_validators = (Arg.string("lhs"), Arg.present("rhs"))
|
||||
|
||||
def handle(self, obj, getter):
|
||||
key, val = self.safe_get_all(obj, getter)
|
||||
getter.ctx.set_var(key, val)
|
||||
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 +325,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 +352,13 @@ 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(),
|
||||
}
|
||||
|
||||
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
|
||||
|
||||
+48
-109
@@ -3,10 +3,6 @@ class Empty:
|
||||
return False
|
||||
|
||||
|
||||
class ExpReturn(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class ExpBase:
|
||||
def __bool__(self):
|
||||
return True
|
||||
@@ -25,10 +21,10 @@ class ExpOp(str, ExpBase):
|
||||
__slots__ = ()
|
||||
|
||||
def eval(self, handlers, *args, **kwargs):
|
||||
value = handlers.get_var(self)
|
||||
if value is Empty:
|
||||
h = handlers.get(self)
|
||||
if h is None:
|
||||
raise KeyError(f"No handler for op/var {self}")
|
||||
return value
|
||||
return h(handlers, *args, **kwargs)
|
||||
|
||||
|
||||
class ExpBinOp(ExpOp):
|
||||
@@ -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,33 @@ class ExpStatements(ExpBase):
|
||||
|
||||
|
||||
class ExprGetter:
|
||||
def __init__(self, obj, ctx, args, kwargs, *, prepend_args=()):
|
||||
GetterEmpty = Empty
|
||||
# class GetterEmpty:
|
||||
# def __bool__(self):
|
||||
# return False
|
||||
|
||||
def __init__(self, obj, handlers, *args, **kwargs):
|
||||
self.obj = obj
|
||||
self.ctx = ctx
|
||||
self.handlers = handlers
|
||||
self.args = args
|
||||
self.prepend_args = prepend_args
|
||||
self.kwargs = kwargs
|
||||
|
||||
def __call__(self, k, *, default=Empty):
|
||||
def __call__(self, k, *, default=GetterEmpty):
|
||||
obj = self.obj
|
||||
if isinstance(k, str):
|
||||
result = obj.kwargs.get_eval(
|
||||
k, self.ctx, *self.args, default=default, **self.kwargs
|
||||
result = (
|
||||
obj.kwargs.get_eval(
|
||||
k, self.handlers, *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)
|
||||
)
|
||||
if result is Empty:
|
||||
if isinstance(k, str)
|
||||
else obj.args.get_eval(k, self.handlers, *self.args, **self.kwargs)
|
||||
)
|
||||
if result is self.GetterEmpty:
|
||||
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
|
||||
@@ -227,26 +169,24 @@ class ExpFunAp(ExpBase):
|
||||
self.kwargs = kwargs if kwargs is not None else ExpDict()
|
||||
|
||||
def eval(self, handlers, *args, **kwargs):
|
||||
handler = handlers.get_handler(self.name)
|
||||
if handler is Empty:
|
||||
handler = handlers.get(self.name)
|
||||
if handler is None:
|
||||
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 ""
|
||||
return f"<FUNAP {self.name}\n{pad}{self.args.pretty_string(depth + 1)}{kwargs_str}\n{pad[:-2]}>"
|
||||
return f"<FUNAP {self.name}\n{pad}{self.args.pretty_string(depth + 1)}{f", {self.kwargs.pretty_string(depth + 1)}" if self.kwargs else ''}\n{pad[:-2]}>"
|
||||
|
||||
def __repr__(self):
|
||||
kwargs_str = f", {self.kwargs}" if self.kwargs else ""
|
||||
return f"<FUNAP:{self.name}{self.args}{kwargs_str}>"
|
||||
return (
|
||||
f"<FUNAP:{self.name}{self.args}{f", {self.kwargs}" if self.kwargs else ''}>"
|
||||
)
|
||||
|
||||
|
||||
class ExpBoundFunAp(ExpFunAp):
|
||||
@@ -269,13 +209,12 @@ class ExpBoundFunAp(ExpFunAp):
|
||||
|
||||
__all__ = (
|
||||
"ExpBase",
|
||||
"ExpBinOp",
|
||||
"ExpBoundFunAp",
|
||||
"ExpDict",
|
||||
"ExpFunAp",
|
||||
"ExpKV",
|
||||
"ExpMethodAp",
|
||||
"ExpOp",
|
||||
"ExpBinOp",
|
||||
"ExpSym",
|
||||
"ExpTuple",
|
||||
"ExpKV",
|
||||
"ExpDict",
|
||||
"ExpFunAp",
|
||||
"ExpBoundFunAp",
|
||||
)
|
||||
|
||||
+26
-116
@@ -1,13 +1,13 @@
|
||||
import contextlib
|
||||
import functools
|
||||
|
||||
from ..latent import ImageBatch
|
||||
from .types import Empty
|
||||
from .util import torch
|
||||
|
||||
|
||||
class Arg:
|
||||
__slots__ = ("default", "name", "validator")
|
||||
__slots__ = ("name", "default", "validator")
|
||||
|
||||
class Empty:
|
||||
pass
|
||||
|
||||
def __init__(self, name, default=Empty, *, validator=None):
|
||||
self.name = name
|
||||
@@ -18,8 +18,9 @@ class Arg:
|
||||
return self.validate(value, *args, **kwargs)
|
||||
|
||||
def validate(self, value):
|
||||
if value is Empty:
|
||||
if self.default is Empty:
|
||||
# FIXME: This shouldn't be using None.
|
||||
if value is None:
|
||||
if self.default is self.Empty:
|
||||
raise ValueError(f"Missing value for argument {self.name}")
|
||||
return self.default
|
||||
try:
|
||||
@@ -31,10 +32,6 @@ class Arg:
|
||||
def tensor(cls, name):
|
||||
return cls(name, validator=ValidateArg.validate_tensor)
|
||||
|
||||
@classmethod
|
||||
def image(cls, name):
|
||||
return cls(name, validator=ValidateArg.validate_image)
|
||||
|
||||
@classmethod
|
||||
def numeric(cls, name, default=Empty):
|
||||
return cls(name, default=default, validator=ValidateArg.validate_numeric)
|
||||
@@ -53,27 +50,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 +60,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 +69,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 +92,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 +129,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):
|
||||
@@ -194,14 +141,6 @@ class ValidateArg:
|
||||
raise ValidateError(f"Expected tensor argument at {idx}, got {type(val)}")
|
||||
return val
|
||||
|
||||
@staticmethod
|
||||
def validate_image(idx, val):
|
||||
if not isinstance(val, ImageBatch):
|
||||
raise ValidateError(
|
||||
f"Expected PIL Image argument at {idx}, got {type(val)}"
|
||||
)
|
||||
return val
|
||||
|
||||
@staticmethod
|
||||
def validate_sequence(idx, val, *, item_validator=None):
|
||||
if not isinstance(val, (list, tuple)):
|
||||
@@ -211,29 +150,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 +158,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 +179,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
|
||||
|
||||
+401
-852
File diff suppressed because it is too large
Load Diff
+10
-112
@@ -1,120 +1,18 @@
|
||||
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):
|
||||
sonar = importlib.import_module("custom_nodes.ComfyUI-sonar")
|
||||
MODULES["sonar"] = sonar.py
|
||||
|
||||
|
||||
__all__ = ("MODULES",)
|
||||
|
||||
+145
-133
@@ -10,43 +10,59 @@ 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
|
||||
|
||||
class FilterHandlerCollection:
|
||||
def __init__(self, op_handlers, refs):
|
||||
self.op_handlers = op_handlers
|
||||
self.refs = refs
|
||||
|
||||
def get(self, k, default=None):
|
||||
result = self.op_handlers.get(k)
|
||||
if result is not None:
|
||||
return result
|
||||
result = self.refs.get(k)
|
||||
if result is None:
|
||||
return None
|
||||
|
||||
def h(obj, *_args, __k=k, __val=result, **_kwargs):
|
||||
if not isinstance(obj, FilterHandlerCollection):
|
||||
raise ValueError(f"Unexpected arguments to variable reference {k}")
|
||||
return result
|
||||
|
||||
return h
|
||||
|
||||
def __contains__(self, k):
|
||||
return k in self.op_handlers or k in self.refs
|
||||
|
||||
def clone_with_refs(self, refs):
|
||||
return self.__class__(self.op_handlers, refs)
|
||||
|
||||
|
||||
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
|
||||
FILTER_HANDLERS = FilterHandlerCollection(
|
||||
expr.BASIC_HANDLERS | expression_handlers.HANDLERS, {}
|
||||
)
|
||||
|
||||
|
||||
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 +74,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 +98,25 @@ 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_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
|
||||
@@ -234,7 +240,7 @@ class Filter:
|
||||
if self.when is None:
|
||||
return True
|
||||
refs = fallback(refs, FilterRefs())
|
||||
matched = self.when.eval(FILTER_HANDLERS.clone(constants=refs, variables={}))
|
||||
matched = self.when.eval(FILTER_HANDLERS.clone_with_refs(refs))
|
||||
# if matched:
|
||||
# print("\nMATCH", self.name)
|
||||
return matched
|
||||
@@ -244,13 +250,7 @@ class Filter:
|
||||
return ops.apply(default_ref, ops, refs=refs)
|
||||
drefs = FilterRefs({"default": default_ref})
|
||||
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}>"
|
||||
return ops.eval(FILTER_HANDLERS.clone_with_refs(refs))
|
||||
|
||||
|
||||
class SimpleFilter(Filter):
|
||||
@@ -419,88 +419,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):
|
||||
|
||||
+47
-579
@@ -1,71 +1,37 @@
|
||||
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
|
||||
from comfy.taesd.taesd import TAESD
|
||||
|
||||
from comfy.utils import bislerp
|
||||
|
||||
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 = (
|
||||
latent.amin(dim=dim, keepdim=True),
|
||||
latent.amax(dim=dim, keepdim=True),
|
||||
)
|
||||
normalized = (latent - min_val).div_(max_val - min_val)
|
||||
return (
|
||||
normalized.mul_(target_max - target_min)
|
||||
.add_(target_min)
|
||||
.clamp_(target_min, target_max)
|
||||
)
|
||||
|
||||
|
||||
# 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,530 +83,32 @@ 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):
|
||||
__slots__ = ()
|
||||
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)
|
||||
|
||||
|
||||
class OCSTAESD:
|
||||
latent_formats = {
|
||||
"sd15": latent_formats.SD15(),
|
||||
"sdxl": latent_formats.SDXL(),
|
||||
}
|
||||
if "sonar" in EXT:
|
||||
get_noise_sampler = EXT["sonar"].noise.get_noise_sampler
|
||||
else:
|
||||
|
||||
@classmethod
|
||||
def get_decoder_name(cls, fmt):
|
||||
return cls.latent_formats[fmt].taesd_decoder_name
|
||||
|
||||
@classmethod
|
||||
def get_encoder_name(cls, fmt):
|
||||
result = cls.get_decoder_name(fmt)
|
||||
if not result.endswith("_decoder"):
|
||||
raise RuntimeError(
|
||||
f"Could not determine TAESD encoder name from {result!r}"
|
||||
)
|
||||
return f"{result[:-7]}encoder"
|
||||
|
||||
@classmethod
|
||||
def get_taesd_path(cls, name):
|
||||
taesd_path = next(
|
||||
(
|
||||
fn
|
||||
for fn in folder_paths.get_filename_list("vae_approx")
|
||||
if fn.startswith(name)
|
||||
),
|
||||
"",
|
||||
)
|
||||
if taesd_path == "":
|
||||
raise RuntimeError(f"Could not get TAESD path for {name!r}")
|
||||
return folder_paths.get_full_path("vae_approx", taesd_path)
|
||||
|
||||
@classmethod
|
||||
def decode(cls, fmt, latent):
|
||||
latent_format = cls.latent_formats[fmt]
|
||||
filename = cls.get_taesd_path(cls.get_decoder_name(fmt))
|
||||
model = TAESD(
|
||||
decoder_path=filename, latent_channels=latent_format.latent_channels
|
||||
).to(latent.device)
|
||||
result = model.decode(latent).movedim(1, 3)
|
||||
return ImageBatch(
|
||||
latent_preview.preview_to_image(result[batch_idx])
|
||||
for batch_idx in range(result.shape[0])
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def img_to_encoder_input(imgbatch):
|
||||
return torch.stack(
|
||||
tuple(
|
||||
torch.tensor(np.array(img), dtype=torch.float32)
|
||||
.div_(127)
|
||||
.sub_(1.0)
|
||||
.clamp_(-1, 1)
|
||||
for img in imgbatch
|
||||
),
|
||||
dim=0,
|
||||
).movedim(-1, 1)
|
||||
|
||||
@classmethod
|
||||
def encode(cls, fmt, imgbatch, latent, *, normalize_output=False):
|
||||
latent_format = cls.latent_formats[fmt]
|
||||
rv = latent_format.process_out(1.0)
|
||||
filename = cls.get_taesd_path(cls.get_encoder_name(fmt))
|
||||
model = TAESD(
|
||||
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))
|
||||
return result.to(latent.dtype).clamp(-rv, rv)
|
||||
|
||||
|
||||
bleh_scale_samples = None
|
||||
UPSCALE_METHODS = ("bicubic", "bislerp", "bilinear", "nearest-exact", "area")
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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,
|
||||
*,
|
||||
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",
|
||||
)
|
||||
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
|
||||
)
|
||||
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)
|
||||
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)
|
||||
|
||||
+75
-168
@@ -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"):
|
||||
filt = filtargs.pop(key, None)
|
||||
if filt is None:
|
||||
continue
|
||||
@@ -167,27 +122,17 @@ 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()
|
||||
@@ -195,7 +140,7 @@ class OCSModel:
|
||||
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
|
||||
|
||||
+64
-746
File diff suppressed because it is too large
Load Diff
+44
-187
@@ -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,
|
||||
batch_size=32,
|
||||
caching=True,
|
||||
cache_reset_interval=9999,
|
||||
set_seed=True,
|
||||
seed_offset=1,
|
||||
cache_reset_interval=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
-115
@@ -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,67 +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)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<Restart: s_noise={self.s_noise:.04}, immiscible={self.immiscible}>"
|
||||
result = result.item()
|
||||
return result * self.s_noise
|
||||
|
||||
@classmethod
|
||||
def simple_schedule(cls, sigmas, start_step, schedule=(), max_iter=1000):
|
||||
@@ -170,27 +74,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)
|
||||
|
||||
+30
-49
@@ -1,20 +1,20 @@
|
||||
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:
|
||||
handlers = None
|
||||
for merge_sampler in merge_samplers:
|
||||
if merge_sampler.when is not None and handlers is None:
|
||||
handlers = FILTER_HANDLERS.clone(constants=ss.refs)
|
||||
# handlers = FILTER_HANDLERS.clone_with_refs(ss.refs)
|
||||
handlers = FILTER_HANDLERS.clone_with_refs(ss.refs)
|
||||
if merge_sampler.check_match(handlers, ss=ss):
|
||||
return merge_sampler
|
||||
return None
|
||||
@@ -43,13 +43,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 +62,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,50 +74,40 @@ 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:
|
||||
nsc.reset_cache()
|
||||
nsc.update_x(x)
|
||||
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:
|
||||
|
||||
+2157
File diff suppressed because it is too large
Load Diff
@@ -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")
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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}
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
+133
-373
@@ -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
|
||||
@@ -35,11 +32,6 @@ class MergeSubstepsSampler:
|
||||
post_filter = options.pop("post_filter", None)
|
||||
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 +52,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):
|
||||
@@ -122,22 +98,6 @@ class MergeSubstepsSampler:
|
||||
def reset(self):
|
||||
pass
|
||||
|
||||
def callback(self, *, ss=None, mr=None, preview_mode=None):
|
||||
ss = fallback(ss, self.ss)
|
||||
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 +112,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))
|
||||
ss.callback()
|
||||
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 +142,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)
|
||||
ss.callback()
|
||||
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 +221,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 +283,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 +326,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,
|
||||
# )
|
||||
@@ -376,9 +356,9 @@ class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
class DivideMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
name = "divide"
|
||||
|
||||
def __init__(self, ss, group, **kwargs):
|
||||
def __init__(self, ss, group, *, schedule_multiplier=4, **kwargs):
|
||||
super().__init__(ss, group, **kwargs)
|
||||
self.schedule_multiplier = self.options.pop("schedule_multiplier", 4)
|
||||
self.schedule_multiplier = schedule_multiplier
|
||||
|
||||
def make_schedule(self, ss):
|
||||
max_steps = len(self.ss.sigmas) - 1
|
||||
@@ -407,24 +387,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:
|
||||
subss.callback()
|
||||
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
|
||||
|
||||
@@ -436,21 +426,19 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
self,
|
||||
ss,
|
||||
group,
|
||||
*,
|
||||
overshoot_expand_steps=1,
|
||||
restart_custom_noise=None,
|
||||
restart=None,
|
||||
**kwargs,
|
||||
):
|
||||
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.overshoot_expand_steps = overshoot_expand_steps
|
||||
restart = fallback(restart, {})
|
||||
self.restart = Restart(
|
||||
s_noise=restart.get("s_noise", 1.0),
|
||||
custom_noise=restart_custom_noise,
|
||||
immiscible=restart.get("immiscible", False),
|
||||
is_flow=ss.model.is_rectified_flow,
|
||||
)
|
||||
|
||||
def make_schedule(self, ss):
|
||||
@@ -480,285 +468,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)
|
||||
subss.callback()
|
||||
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,
|
||||
}
|
||||
|
||||
+17
-175
@@ -1,18 +1,9 @@
|
||||
from typing import NamedTuple
|
||||
|
||||
import torch
|
||||
|
||||
from comfy.k_diffusion.sampling import get_ancestral_step
|
||||
|
||||
from .filtering import FilterRefs
|
||||
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:
|
||||
@@ -90,7 +81,6 @@ class StepSamplerGroups(CommonOptionsItems):
|
||||
|
||||
class SamplerState:
|
||||
CLONE_KEYS = (
|
||||
"cfg_scale_override",
|
||||
"model",
|
||||
"hist",
|
||||
"extra_args",
|
||||
@@ -132,7 +122,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 +137,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 +151,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 +159,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 +173,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__)
|
||||
@@ -325,42 +191,18 @@ class SamplerState:
|
||||
obj.update()
|
||||
return obj
|
||||
|
||||
def callback(self, hi=None, *, preview_mode="denoised"):
|
||||
def callback(self, hi=None):
|
||||
if not self.callback_:
|
||||
return None
|
||||
hi = self.hcur if hi is None else hi
|
||||
if preview_mode == "cond":
|
||||
preview = fallback(hi.denoised_cond, hi.denoised)
|
||||
elif preview_mode == "uncond":
|
||||
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": hi.denoised,
|
||||
})
|
||||
|
||||
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)
|
||||
|
||||
@@ -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
|
||||
@@ -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",
|
||||
))
|
||||
+12
-107
@@ -1,119 +1,24 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
|
||||
import torch
|
||||
from comfy.k_diffusion.sampling import to_d
|
||||
|
||||
F = torch.nn.functional
|
||||
from comfy.k_diffusion.sampling import to_d
|
||||
|
||||
|
||||
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))
|
||||
|
||||
|
||||
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
|
||||
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 noise.sub_(mean).div_(std).mul_(factor)
|
||||
|
||||
|
||||
def find_first_unsorted(tensor, desc=True):
|
||||
|
||||
@@ -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])),
|
||||
)
|
||||
Reference in New Issue
Block a user