Refactor (#1)

Refactor all the things!
This commit is contained in:
blepping
2024-08-04 11:28:06 -06:00
committed by GitHub
parent 8d0716623c
commit 470b38231f
28 changed files with 7237 additions and 1162 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
MIT License
Copyright (c) 2024 blepping
Copyright (c) 2024 blepping <https://github.com/blepping>
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
+568 -77
View File
@@ -1,93 +1,584 @@
# Overly Complicated Sampling
Wildly unsound and experimental sampling for [ComfyUI](https://github.com/comfyanonymous/ComfyUI).
## Description
Experimental and mathematically unsound (but fun!) sampling for [ComfyUI](https://github.com/comfyanonymous/ComfyUI).
Very unstable, experimental and mathematically unsound sampling for ComfyUI.
**Status**: In flux, may be useful but likely to change/break workflows frequently. Mainly for advanced users.
Current status: In flux, not suitable for general use.
*Note*: You will basically always have to tweak settings like `s_noise` to get a good result. If the generation looks smooth/undetailed increase `s_noise` somewhere. If it looks crunchy, super high contrast, etc then try reducing noise.
## Features
## Nodes
* Many different samplers.
* Allows scheduling samplers (i.e. run `euler` for steps 1-4, then switch to `dpmpp_sde`).
* CFG++ support (for some samplers).
* Native support for Restart sigmas.
* Supports custom noise types.
* Immiscible noise for sampling and Restart. See https://arxiv.org/abs/2406.12303 (note that it was designed for training not inference).
* Allows splitting/combining steps in various ways for (potentially) more accurate sampling.
* Supports Diffrax, torchdiffeq, torchode and torchsde solver backends. (SDE mode not recommended currently.)
* Many tuneable parameters to play with.
### ComposableSampler
**Possible Parameters**
* `avgmerge_stretch`(`0.4`): Used for `average` and `sample` merge types. See below.
* `model_call_cache`(unset): Caches the result of model calls. For example, Bogacki is 3 model calls per step. The first one usually depends on the merge strategy: `average` for example shares the first model evaluation between substeps, but subsequent model calls (i.e. Bogacki 2nd and 3rd model evaluations) still occur. When the model call cache is active, it's possible to cache those evalutions and avoid a model call for the remaining substeps. If you set `model_call_cache` to `1` then the result of that second call will be cached and if you're running two Bogacki substeps then the second one will use the cached version. Massively accelerates inference (especially when using the `average` merge strategy) but is likely very unsound and inaccurate. Does not apply to the sampler call for the `sample` merge strategy.
* `model_call_cache_threshold`(`1`): Disables caching model call results with a call index below the threshold value (starting at 0). For example, if set to `2` and using a sampler like Bogacki that calls the model two extra times, the first will never be cached. The default value of `1` disabling caching for the first model call per substep. I generally would not recommend setting it to `0`, especially with `average` or `sample` merge strategies.
* `model_call_cache_max_use`(`1000000`): The number of times cache items can be re-used. The default is effectively no limit. Where would this be useful? Let's say you're using the `average` merge strategy and a multi step sampler that calls the model at least one more time with 50 substeps. If you set the value to `25`, the model cache result will be updated around substep 25 which _may_ produce better results than reusing the result 50 times.
Since it's kind of confusing even for me, a little more explanation: The model call cache caches results for model call indexes between `model_call_cache_threshold` and `model_call_cache - 1`. If you set `model_call_cache_threshold` to `0` and `model_call_cache` to `1` then only the first model call will be cached. If you set `model_call_cache_threshold` to `1` and `model_call_cache` to `2` then call 0 will not be cached, call 1 will be cached, call 2 will be cached, call 3 will not be cached, and so on.
#### Merging
When running multiple substeps per step, the results will combined based on the merge strategy. Possible strategies (in order of least weird to most weird):
* `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.
* `normal`: The model is called at least once per substep (and possibly additional times for higher order samplers). The result of each substep is noised and the next substep uses that result. Then all the results are averaged.
* `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.
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.
### ComposableStepSampler
This node has a text input for YAML (or JSON) advanced parameters.
For example, you could enter something like this in the field:
```yaml
reta: 1.1
leap: 3
dyn_deta_mode: "deta"
```
**Possible Parameters**
#### General
* `eta`(`1.0`): Will override `eta` in the node if set.
* `dyn_eta_start`(`unset`) and `dyn_eta_end`(`unset`): No effect unless both values are set. Will interpolate between start and end based on the percentage of sampling. *Note*: This is a factor applied to ETA, not a flat value.
* `s_noise`(`1.0`): Will override `s_noise` in the node if set.
* `solver_type`(`midpoint`): Applies to DPM++ 2m SDE. May be one of `midpoint` or `heun` (`midpoint` is generally recommended).
#### Reversible
* `reta`(`1.0`): Reverse ETA.
* `dyn_reta_start`(`unset`) and `dyn_reta_end`(`unset`): No effect unless both values are set. Will interpolate between start and end based on the percentage of sampling. *Note*: This is a factor applied to RETA, not a flat value.
#### Dancing
* `leap`(`2`): Distance to try to leap forward. If you set `leap` to `1` you just get plain old Euler ancestral.
* `deta`(`1.0`): ETA used for dance steps.
* `dyn_deta_start`(`unset`) and `dyn_deta_end`(`unset`): No effect unless both values are set. Will interpolate between start and end based on the percentage of sampling. *Note*: This is a factor applied to DETA, not a flat value.
* `dyn_deta_mode`(`lerp`): May be one of:
* `deta`: Scales `deta` based on the value from `dyn_deta_start/end`.
* `lerp`: Does the dance step according to `deta` and then LERPs the non-dance sample result with the dance sample result based on the scale calculated from `dyn_deta_start/end` (which is `1.0` if they are unset). For example, if the dance scale is `0.5` you will get 50% normal sampling, 50% dancing sampling.
* `lerp_alt`: Similar to `lerp` except it LERPs with the leap result instead of a normal Euler ancestral result.
#### RES
* `res_simple_phi`(`false`): Uses a faster but possibly less accurate method for calculating phi. What does phi do? I haven't the foggiest!
* `res_c2`(`0.5`): Solver partial step size, the default of `0.5` appears to use the midpoint. Setting it to a lower value might possibly be more accurate but slower?
#### TTM JVP
`alterate_phi_2_calc`(`true`): Supposedly works better than disabled when ETA isn't 0. I didn't notice a difference.
**Note**: TTM is a weird sampler. If you're using model caching you must make sure the entries TTM uses are populated first (by having before any other samplers that call the model multiple times). It may also not work with some other model patches and upscale methods.
## Credits
I can move code around but sampling math and creating samplers is far beyond my ability. I didn't write any of the original samplers:
* Euler, DPMPP SDE, DPMPP 2S, DPM++ 2m, 2m SDE and 3m SDE samplers based on ComfyUI's implementation.
* 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).
* 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).
* Normal substep merge strategy based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
* Immiscible noise processing based on implementation from https://github.com/kohya-ss/sd-scripts/pull/1395 and idea for sampling with it from https://github.com/Clybius
* Precedence climbing (Pratt) expression parser based on implementation from https://github.com/andychu/pratt-parsing-demo
Thanks!
This repo wouldn't be possible without building on the work of others. Thanks!
## Usage
First, a note on the basic structure:
![Basic nodes example](assets/basic_sampling.png)
The sampler node connects to a group node. You can chain group nodes, however only one can match per step. Checking for a match starts at the group furthest from the sampler node. So if you have `Group1 -> Group2 -> Group3 -> Sampler`, the groups will be tried starting from `Group1`. Currently matching groups is only time based.
You will then connect a substeps node to the group. These can also be chained and like groups, execution starts with the node furthest from the group. I.E.: `Substeps1 -> Substeps2 -> Substeps3 -> Group` will start with `Substeps1`.
Most of these nodes have a text parameter input (YAML format - JSON is also valid YAML so you can use that instead if you prefer) and a parameter input. The parameter input can be used to specify stuff like custom noise types.
You may use filters and expressions in the text parameter input. See:
* [Filters](docs/filter.md)
* [Expressions](docs/expression.md)
## Nodes
### `OCS Sampler`
The main sampler node, with an output suitable for connecting to a `SamplerCustom`. This node has builtin support for Restart sampling, if you are
using Restart don't use the `RestartSampler` node.
You can connect a chain of `OCS Group` nodes to it and it will choose one per step (based on conditions like time).
#### Input Parameters
* `restart_custom_noise`: Value type: `SONAR_CUSTOM_NOISE`. Allows specifying a custom noise type when used with Restart sampling.
#### Text Parameters
Shown in YAML with default values.
<details>
<summary>★★ Expand ★★</summary>
```yaml
# Noise scale. May not do anything currently.
s_noise: 1.0
# ETA (basically ancestralness). May not do anything currently.
eta: 1.0
# Reversible ETA (used for reversible samplers). May not do anything currently.
reta: 1.0
# Parameters related to restart sampling.
restart:
# Scales the noise added by restart sampling.
s_noise: 1.0
# Immiscible block same as described below.
immiscible:
size: 0
# The noise block allows defining global noise sampling parameters.
noise:
# You can disable this to allow GPU noise generation. I believe it only makes a difference for Brownian.
cpu_noise: true
# ComfyUI has a bug where if you disable add_noise in the sampler, no seed gets set. If you
# are manually noising a sample and have add_noise turned off then you should enable this if
# you want reproducible generations.
set_seed: false
# Global scale scale for generated noise
scale: 1.0
# Whether the generated noise should be normalized before use. Generally a good idea to leave enabled.
normalize_noise: true
# Dimensions to normalize over (when normalization is enabled). Negative values mean starting
# from the end (i.e. -1 means the last dimension, -2 means the penultimate dimension).
# Latents generally have these dimensions: batch, channels, height, width
# The default of [-3, -2, -1] normalizes noise over the batch. You can try something like
# [-2, -1] to normalize over the batch and channels. See: https://pytorch.org/docs/stable/generated/torch.std.html
normalize_dims: [-3, -2, -1]
# 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: 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 you will generally want to reset each step.
cache_reset_interval: 1
# Immiscible noise processing, see: https://arxiv.org/abs/2406.12303
immiscible:
# Batch size, 0 disables.
size: 0
# Reference mode, values can be one of:
# x: Uses the current latent as a reference.
# noise: Uses the current noise as a reference (x - denoised)
# denoised: Uses the model image prediction as a reference (factors in positive and negative prompts).
# uncond: The model unconditional prediction (negative prompt)
# cond: The model conditional prediction (positive prompt)
# Advanced feature: Additionally you may enter a string of operations in the format:
# "x - denoised * 2 + cond" (just an example, not a recommended setting)
# Possible operations: + - / * min max add sub div mul
# Note: Each value and operation must be space delimited (i.e. "x-1" will not work).
# Also normal operator precedence does not apply here.
ref: default
# Batching mode, one of:
# batch: Matches vs batches. Immiscible mode is disabled if size < 2
# channel: Splits the batch into a list of channels and matches against those.
# row: Splits the batch into a list of rows and matches against those.
# column: Splits the batch into a list of columns and matches against those.
# Note: Requires reshaping both the noise and x, may be slow and consume
# a lot of VRAM.
batching: channel
# Scale for reference latent. Can be negative.
scale_ref: 1.0
# Allows normalizing the reference. If this is a list, you can specify the dimensions to
# normalize. See normalize_dims above and https://pytorch.org/docs/stable/generated/torch.std.html
normalize_ref: false
# The proportion of immiscible-ized noise.
# You get (immiscible_noise * strength) + ((1.0 - strength) * normal_noise) - LERP.
strength: 1.0
# See: https://docs.scipy.org/doc/scipy/reference/generated/scipy.optimize.linear_sum_assignment.html#scipy.optimize.linear_sum_assignment
maximize: false
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:
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: null
denoised: null
jdenoised: null
cond: null
uncond: null
```
</details><br/>
Any parameters you don't specify will use the defaults. For example if your text parameter block is:
```yaml
noise:
cpu_noise: false
```
Then the rest of the parameters will use the defaults shown above.
### `OCS Group`
Defines a group of substeps.
#### Merging
When running multiple substeps per step, the results will combined based on the merge strategy. Possible strategies (in order of least weird to most weird):
* `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.
* `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.
<!--
* `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.
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
* `merge_method`: One of the merge methods described above in the Merging section.
* `time_mode`(`step`): One of `step`, `step_pct`, `sigma`. Time matching mode. Matching based on steps generally will be simplest. Matches are inclusive and steps start at 0 (so step 0 is the first step). `step_pct` is the percentage of total steps (1.0=100%, 0.5=50%, etc).
* `time_start`(`0`): Match start time.
* `time_end`(`999`): Match end time.
Example:
![Group time filter example](assets/group_time_example.png)
The left side group matches steps 0, 1, 2. The right side group matches all steps. This setup will use whatever substeps are connected to the first group for the first three steps and the second group will handle the rest.
#### Input Parameters
<!--
* `merge_sampler`: Value type: `OCS_SUBSTEPS`. Only used when `merge_method` is `sample` or `sample_uncached`. Allows defining the sampler used for merging substeps.
-->
* `restart_custom_noise`: Currently only used by the `overshoot` merge method.
#### Text Parameters
Shown in YAML with default values.
<details>
<summary>★★ Expand ★★</summary>
```yaml
# Noise scale. May not do anything currently.
s_noise: 1.0
# ETA (basically ancestralness). May not do anything currently.
eta: 1.0
# Reversible ETA (used for reversible samplers). May not do anything currently.
reta: 1.0
# Expression.
when: null
# Interpolate the schedule by the specified factor. Only used by the overshoot merge method.
#: Example if factor 2 and steps [0,1,2] you'd get [0, 0.5, 1.0, 1.5, 2]
overshoot_expand_steps: 1
# Only used by the overshoot merge method currently.
restart:
# Scales the noise added by restart sampling.
s_noise: 1.0
# Immiscible block same as described above.
immiscible:
size: 0
pre_filter: null
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.)
* `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_sde`
* `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.
* `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`:
* `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**
|Name|Cost|History|Order|Reversible|CFG++|
|-|-|-|-|-|-|
|`adapter`|?|?|?|?|?|
|`bogacki`|2|||||
|`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|||||
|`euler_cycle`|1||||X|
|`euler_dancing`|1|||||
|`euler`|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)||||
|`res`|2|||||
|`reversible_bogacki`|2|||X||
|`reversible_heun`|2|||X||
|`reversible_heun_1s`|1|1||X||
|`rk4`|4|||||
|`solver_diffrax`|variable|||||
|`solver_torchdiffeq`|variable|||||
|`solver_torchode`|variable|||||
|`solver_torchsde`|variable|||||
|`trapezoidal`|2|||||
|`trapezoidal_cycle`|2|||||
|`ttm_jvp`|2|||||
`deis`, `ipndm*` do not seem to work well with ancestralness, I recommend `eta: 0.25` or disable it completely.
**Solver Backend Samplers**:
You will need to have the relevant Python package installed in your venv to use these. TDE cannot handle batches and
each batch item will be evaluated separately. Using `tode` may be faster for batch sizes over 1.
`ode_solver` types for TDE: adaptive: `dopri8`, `dopri5`, `bosh3`, `fehlberg2`, `adaptive_heun`, fixed step: `euler`, `midpoint`, `rk4`, `explicit_adams`, `implicit_adams`
`ode_solver` types for TODE: adaptive only: `dopri5`, `tsit5`, `heun`. I haven't much luck with anything other than `dopri5`.
Note that adaptive solvers may be _very_ slow. Think along the lines of 20-100 model calls per substep (or in other words, the equivalent for running that many `euler` steps). Tolerances only apply to adaptive solvers.
**Cycle Samplers** (`euler_cycle`, `trapezoidal_cycle`)
Basically a different approach to ancestral sampling. First a crash course on how sampling works:
Each step has an expected noise level, with the first step generally being pure noise and the end of the last step aiming to end with no noise remaining. Let's say the image on the current step is called `x`, calling the model with `x` gives us a prediction of what the image looks like with all noise removed (`denoised`), however the model is not capable of just removing all the noise in a single step: its prediction will be imprecise. `x - denoised` leaves us with just the noise (we subtract the prediction which theoretically has no noise from the noisy sample). This is a very simplified, but the idea is basically to add the noise back into `denoised`, but scaled so that it matches the amount of noise expected on the _next_ step. `denoised + noise * expected_noise_at_next_step`.
When doing ancestral sampling, we actually _overshoot_ expected noise for the next step and add less than that amount back to `denoised`. Then we generate some of our own noise and add it, scaled so that the result matches `expected_noise_at_next_step`. `eta` controls how the scale of the overshoot.
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.
#### 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.
* `step_method`(`euler`): Method used for sampling the substeps. May include a parenthesized number (i.e. `rk4 (3)`) which denotes the number of _extra_ model calls required per sample. At least one is always required. So `euler` requires 1 in total, `rk4` requires 4 in total. RK4 is about 4 times slower than `euler`.
#### Input Parameters
* `custom_noise`: Value type: `SONAR_CUSTOM_NOISE`. Allows specifying a custom noise type for samplers that generate noise (most of them).
#### Text Parameters
Shown in YAML with default values.
<details>
<summary>★★ Expand ★★</summary>
```yaml
# Scale for added noise.
s_noise: 1.0
# ETA (basically ancestralness).
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, eta*dyn_eta_start at the beginning,
# eta*dyn_eta_end at the end.
dyn_eta_start: null
dyn_eta_end: null
# alt CFG++ scale (see https://cfgpp-diffusion.github.io/)
# Based on the initial incorrect ComfyUI implementation, but it seems to
# produce decent results sometimes.
# Can also be set to a negative value (I don't recommend going lower than -0.5).
alt_cfgpp_scale: 0
# CFG++ (see https://cfgpp-diffusion.github.io/)
cfgpp: false
### Reversible Settings ###
# 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
post_filter: null
### ODE Sampler Settings ###
# Solver type.
de_solver: dopri5 # Example - varies based on solver sampler.
# Relative tolerance (log 10)
de_rtol: -1.5
# Absolute tolerance (log 10)
de_atol: -3.5
# Max model calls allowed to compute the solution. If the limit is exceeded, it is an error.
de_max_nfe: 1000
# Min sigma to sample to. If the current step start <= min sigma, then the sampler will run
# a Euler step. If the current step end <= min sigma then the slover will sample to the min
# sigma and then to a Euler step from min sigma for the rest.
de_min_sigma: 0.0292
# Hack that seems to help results by stretching the down sigma a bit. Set to 0 to disable.
de_fixup_hack: 0.025
# Used to split the step into sections. Useful for fixed step methods.
# Applies to: solver_torchode, solver_diffrax
de_split: 1
# Initial step size (as a percentage).
# Applies to: solver_torchode, solver_diffrax
de_initial_step: 0.25
# Coefficients for the step size PID controller.
# See https://en.wikipedia.org/wiki/Proportional%E2%80%93integral%E2%80%93derivative_controller
# These values seem okay with dopri5.
# Applies to: solver_tode, solver_diffrax
de_ctl_pcoeff: 0.3
de_ctl_icoeff: 0.9
de_ctl_dcoeff: 0.2
# Controls whether to compile the solver. May or may not work,
# also may or may not be a speed increase as the compiled solver is
# not cached between substeps.
# Applies to: solver_torchode
tode_compile: false
### torchsde solver specific parameters.
tsde_noise_type: "scalar"
tsde_sde_type: "stratonovich"
tsde_levy_area_approx: "none"
tsde_noise_channels: 1
tsde_g_multiplier: 0.05
tsde_g_reverse_time: true
tsde_g_derp_mode: false
tsde_batch_channels: true
### diffrax solver specific parameters.
# Turns on adaptive stepping. When enabled, de_split is not used.
# When disabled, it may be desirable to set de_split.
diffrax_adaptive: false
# Hack to make some solver methods work. May not be safe.
diffrax_fake_pure_callback: true
# Some diffrax methods don't allow adaptive stepping, enabling this
# makes them usable although it's less efficient (3x cost, 2x accuracy).
diffrax_half_solver: false
diffrax_batch_channels: false
# Some solvers require specific types of Levy area approximation.
# See: https://docs.kidger.site/diffrax/api/brownian/#levy-areas
diffrax_levy_area_approx: "brownian_increment"
# Some solvers may require manually specifying the error order.
diffrax_error_order: null
# Enables SDE mode (and SDE-specific solvers). May not be worth using.
diffrax_sde_mode: false
# Noise multiplier when SDE mode is enabled.
diffrax_g_multiplier: 0.0
# Only applies when time scaling is enabled. Reverses time.
diffrax_g_reverse_time: false
# Scales the g multiplier based on the current time.
diffrax_g_time_scaling: false
# Experimental option to flip the sign on the g multiplier when time >= half the step.
# i.e. if you'd get 1,2,3,4 as g values for the step, with this it would be 1,2,-3,-4.
diffrax_g_split_time_mode: false
### Other Sampler Specific Parameters ###
# Used for some samplers that use history from previous steps.
# List of samplers and default value below:
# dpmpp_2m: 1
# dpmpp_2m_sde: 1
# dpmpp_3m_sde: 2
# reversible_heun_1s: 1
# ipndm: 1 (max 3)
# ipndm_v: 1 (max 3)
# deis: 1 (max 3)
history_limit: 999 # Varies based on sampler.
# Used for some samplers with variable order. List of samplers and default value below:
# heunpp2: 3
max_order: 999 # Varies based on sampler.
# Used for dpmpp_2m. One of midpoint, heun
solver_type: "midpoint"
# Coefficients mode for DEIS. One of tab or rhoab.
deis_mode: "tab"
# Used for samplers with cycle in the name. Controls how much noise is cycled per step.
cycle_pct: 0.25
# Used for ttm_jvp. Supposed works better when ETA > 0
alternate_phi_2_calc: true
# Parameters for dancing samplers:
# Number of steps to leap ahead.
leap: 2
# ETA for dance steps
deta: 1.0
# dyn_deta works the same as dyn_eta/reta. See above.
dyn_deta_start: null
dyn_deta_end: null
# One of lerp, lerp_alt, deta
dyn_deta_mode: "lerp"
```
</details>
### `OCS SimpleRestartSchedule`
Generates a restart schedule.
#### Node Parameters
* `start_step`: 0-based first step for the restart schedule to apply.
#### Input Parameters
* `sigmas`: Sigmas to restartify. Output from any normal schedule node.
#### Text Parameters
JSON or YAML schedule in list form.
```yaml
- [4, -3]
- [2, -1]
- 1
```
Each item should be one of:
* A pair `[interval, jump]` - after `interval` steps, make a relative jump of `jump` steps.
* A single integer `schedule_index`: resume the schedule at the specified 0-based index.
The example above means:
1. After 4 steps, jump back 3 steps.
2. After 2 steps, jump back one step.
3. Go to the second item (after 2 steps, jump back one step).
The node `start_step` parameter is effectively the same as `[start_step, 0]` as a schedule item.
+7 -2
View File
@@ -2,7 +2,12 @@ from .py import nodes
NODE_CLASS_MAPPINGS = {
"ComposableSampler": nodes.ComposableSampler,
"ComposableStepSampler": nodes.ComposableStepSampler,
"OCS Sampler": nodes.SamplerNode,
"OCS Substeps": nodes.SubstepsNode,
"OCS Group": nodes.GroupNode,
"OCS Param": nodes.ParamNode,
"OCS MultiParam": nodes.MultiParamNode,
"OCS ModelSetMaxSigma": nodes.ModelSetMaxSigmaNode,
"OCS SimpleRestartSchedule": nodes.SimpleRestartSchedule,
}
__all__ = ["NODE_CLASS_MAPPINGS"]
Binary file not shown.

After

Width:  |  Height:  |  Size: 48 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 31 KiB

+139
View File
@@ -0,0 +1,139 @@
# OCS Expressions
See [Filters](filter.md) for places where expressions apply.
## Expressions
OCS implements a simple expression language.
Supported math operators: `+`, `-`, `*`, `/`, `//` (integer division), `**` (power)
Supported logic operators: `||`, `&&`, `==`, `!=`, `>`, `<`, `>=`, `<=`
Operator precedence should generally work the way you'd expect.
You may surround a function name with backticks to turn it into a binary operator (only for functions that take two arguments).
Functions are called via `name(param1, param2)`. Keyword arguments may be passed using the `:>` operator. Example:
`name(param1, key :> 123, key2 :> otherfunction(10))`.
Symbols (simple string type) are defined using `'symbol_name` - note the solitary single quote. They may not contain spaces.
`;` can be used to sequence operations. I.E. `exp1 ; exp2` evaluates `exp1`, then `exp2` and then result of the expression is whatever `exp2` returned.
Like Python, a parenthesized expression with a trailing comma can be used to create an empty tuple. Example: `(1,)`
## Filter Variables
Indexes like `step` are zero-based: `0` will be the first step.
### Basic Variables
* `default`: Context specific default value. i.e. if used in an `input` expression this would be `x`, if used for `output` this would be the current result.
* `step`: Current step.
* `substep`: Current substep.
* `dt`: `sigma_next - sigma`
* `sigma_idx`: Index of the current sigma. Note that when using restarts this will be based on the restart sigma chunks, not full sigma list.
* `sigma`: The current sigma.
* `sigma_next`: The next sigma.
* `sigma_down`: The down sigma in ancestral sampling.
* `sigma_up`: The up sigma in ancestral sampling.
* `sigma_prev`: The previous sigma (may be `None`).
* `hist_len`: Current available history length. "Now" counts as one.
* `sigma_min`: The minimum sigma (based on the full list).
* `sigma_max`: The maximum sigma (based on the full list).
* `step_pct`: Percentage for the current step (based on total steps).
* `total_steps`: Total steps to be sampled.
### Extended Variables
* `denoised`: From the current step or substep. May not be available in model `input` or group `pre_filter`.
* `cond`: From the current step or substep. May not be available in model `input` or group `pre_filter`.
* `uncond`: From the current step or substep. May not be available in model `input` or group `pre_filter`.
* `denoised_prev`: Only available when model history exists.
* `cond_prev`: Only available when model history exists.
* `cond_prev`: Only available when model history exists.
### Model Filter Variables
* `model_call`: Only applicable to `model` filters, will be the model call index. I.E. if the sampler calls the model three times, the filter would be called with model call indexes `0`, `1` and `2`.
Available in model filters, with the exception of the `input` filter.
* `denoised_curr`
* `cond_curr`
* `uncond_curr`
## Basic Expression Functions
| | Name | Input | Output |
| :--- | :--- | :--- | :--- |
|⬤| `all` | `B`\* | `B` |
| <td colspan=3 align=left>Evaluates to true if all its arguments evaluate to true. <br/> **Example:** `all(x > 1, y < 1)`</td> |
|⬤| `any` | `B`\* | `B` |
| <td colspan=3 align=left>Evaluates to true if any of its arguments evaluate to true. <br/> **Example:** `any(x > 1, y < 1)`</td> |
|⬤| `between` | value:`N`, from:`N`, to:`N` | `B` |
| <td colspan=3 align=left>Boolean range checking. <br/> **Example:** `between(value, low, high)`</td> |
|⬤| `comment` | `*` | `null` |
| <td colspan=3 align=left>Ignores any arguments passed to it (they won't be evaluated at all but must parse as a valid expression) and returns `None`</td> |
|⬤| `dict` | `*`* | `dict` |
| <td colspan=3 align=left>Constructs a dictionary from its keyword arguments. _Note_: You may not pass positional arguments. <br/> **Example:** `dict(key1 :> value1, keyN :> valueN)` |
|⬤| `get` | name:`SY`, fallback:`*` | `*` |
| <td colspan=3 align=left>Returns a variable if set, otherwise the fallback. <br/> **Example:** `get('somevar, 123)`</td> |
|⬤| `if` | condition:`B`, then:`*`, else:`*` | `*` |
| <td colspan=3 align=left>Conditional expressions. <br/> **Example:** `if(condition, true_expression, false_expression)`</td> |
|⬤| `index` | index:`IDX`, value:`S \| T` | `*` |
| <td colspan=3 align=left>Index function.</td> |
|⬤| `is_set` | name:`SY` | `B` |
| <td colspan=3 align=left>Tests whether a variable is set.</td> |
|⬤| `max` | values:`SN` | `N` |
| <td colspan=3 align=left>Maximum operation. _Note_: Takes one sequence argument. <br/> **Example:** `min((1, 2, 3))`</td> |
|⬤| `min` | values: `SN` | `N` |
| <td colspan=3 align=left>Minimum operation. _Note_: Takes one sequence argument. <br/> **Example:** `max((1, 2, 3))`</td> |
|⬤| `mod` | lhs:`N`, rhs:`N` | `N` |
| <td colspan=3 align=left>Modulus operation: <br/> **Example:** `mod(5, 2)`</td> |
|⬤| `neg` | `N` | `N` |
| <td colspan=3 align=left>Negation. <br/> **Example:** `neg(2)`</td> |
|⬤| `not` | `B` | `B` |
| <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> |
|⬤| `unsafe_call` | `callable`, `*`\* | `*` |
| <td colspan=3 align=left>Allows calling an arbitrary callable. <br/> **Example:** `unsafe_call(some_callable, arg1, arg2, kwarg1 :> 123)`</td>
**Legend**: `B`=boolean, `N`=numeric, `NS`=scalar numeric, `I`=integer, `F`=float, `T`=tensor, `S`=sequence, `SN`=numeric sequence, `SY`=symbol, `*`=any -- parenthized values indicate argument defaults. `*` following the type indicates variable length arguments. For functions that take keyword arguments, the type will be written like "_name: `TYPE(default_value)`_".
## Tensor Expression Functions
*Tensor dimensions hint*: Most tensors you'll be dealing with are laid out as `batch`, `channels`, `height`, `width`. Negative indexes start from the end, so dimension `-1` would mean _width_ just the same as `3`.
| | Name | Input | Output |
| :--- | :--- | :--- | :--- |
|⬤| `t_bleh_enhance` | tensor:`T`, mode:`SY`, scale:`N(1.0)` | `T`
| <td colspan=3 align=left>Available if you have the [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) node pack installed. See [Filtering](filter.md#bleh_enhance). <br/> **Example:** `bleh_enhance(some_tensor, 'bandpass, 0.5)`</td> |
|⬤| `t_blend` | tensor1:`T`, tensor2:`T`, scale:`N(0.5)`, mode:`SY(lerp)` | `T` |
| <td colspan=3 align=left>Tensor blend operation. <br/> **Example:** `t_blend(t1, t2, 0.75, 'lerp)`</td> |
|⬤| `t_contrast_adaptive_sharpening` | tensor:`T`, scale:`N(0.5)` | `T` |
| <td colspan=3 align=left>Contrast adaptive sharpening. _Note_: Not recommended to call on noisy tensors (so `denoised` but probably not `x`). <br/> **Example:** `t_contrast_adaptive_sharpening(some_tensor, 0.1)`</td> |
|⬤| `t_flip` | tensor:`T`, dim:`NS`, mirror:`B(false)` | `T` |
| <td colspan=3 align=left>Flips a tensor on the specified dimension. If the third argument is true, it will be mirrored around the center in that dimension. <br/> **Example:** `t_flip(some_tensor, -1, true)`</td> |
|⬤| `t_mean` | tensor:`T`, dim:`SN(-3, -2, -1)` | `T` |
| <td colspan=3 align=left>Tensor mean, second argument is dimensions. <br/> **Example:** `t_mean(some_tensor, (-2, -1))`</td> |
|⬤| `t_noise` | tensor:`T`, type:`SY(gaussian)` | `T` |
| <td colspan=3 align=left>Generates un-normalized noise (use `t_norm` if you want to normalize it). If you have ComfyUI-sonar you can use any noise type that supports, otherwise only `gaussian`. The generated noise will have the same shape as the supplied tensor (hopefully, may not be true for every exotic noise type but at least should be broadcastable to the tensor). <br/> Example: `t_noise(some_tensor, 'pyramid)`</td> |
|⬤| `t_norm` | tensor:`T`, factor:`N(1.0)`, dim:`SN(-3, -2, -1)` | `T` |
| <td colspan=3 align=left>Tensor normalization (subtracts mean, divides by std). <br/> **Example:** `t_norm(some_tensor, 1.0, (-2, -1))`</td> |
|⬤| `t_sonar_power_filter` | tensor:`T`, filter:`dict` | `T` |
| <td colspan=3 align=left>Available if you have [ComfyUI-sonar](https://github.com/blepping/ComfyUI-sonar) installed. See [Filtering](filter.md#sonar_power_filter). Constructs a power filter from a dictionary argument. _Note_: May be slow as the filter is reconstructed on every evaluation. <br/> **Example:** `t_sonar_power_filter(some_tensor, dict(alpha :> 0.1, min_freq :> 0.2, max_freq :> 0.6))`</td> |
|⬤| `t_roll` | tensor:`T`, amount:`NS(0.5)`, dim:`SN((-2,))` | `T` |
| <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_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> |
|⬤| `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.
+187
View File
@@ -0,0 +1,187 @@
# OCS Filters
Filters allow changing sampler inputs/outputs, model input/outputs, generated noise and so on. They are
configured by the advanced YAML/JSON parameter block in the node.
## Filter Support
### `OCS Substeps`
Set via the `pre_filter` and `post_filter` keys.
### `OCS Group`
Set via the `pre_filter` and `post_filter` keys.
*Note*: Since the group pre-filter may be called before any model calls, variables like `denoised` may not be available
in expressions.
### `OCS Sampler`
**Noise**
```yaml
noise:
# Or set to a valid filter definition.
filter: null
```
**Model**
*Note*: Since the model filters may be called before any other model calls, variables like `denoised` may not be available
in expressions. With the exception of the `input` filter you will have access to `denoised_curr`, `cond_curr`, etc.
See [Expressions](expression.md#model-filter-variables).
```yaml
model:
filter:
# Applies to the input passed to the model.
input: null
# Applies to denoised output.
denoised: null
# Applies to JVP denoised output (only used by TTM sampler)
jdenoised: null
# Applies to cond output.
cond: null
# Applies to uncond output.
uncond: null
```
### Immiscible
`immiscible` is a special type of filter: in this case, you do not set `filter_type`. Normal filter keys
apply in the places where `immiscible` can be set.
## Filter Definitions
For information about expressions, see [Expressions](expression.md).
A basic filter supports these keys:
```yaml
enabled: true
filter_type: simple
# Expression that is evaluated to determine whether the filter applies. May be null.
# If set, should evaluate to a boolean.
when: null
blend_mode: lerp
# Blend strength applied to output.
strength: 1.0
# Input expression. input, ref, output and final should evaluate to a tensor.
input: default
# Reference expression (only used for immiscible noise currently).
ref: default
# Output expression.
output: default
# Final expression - occurs *after* blending.
final: default
```
There may be additional keys depending on the filter type.
## Filter Types
### `simple`
Base filter, no special behavior. No additional parameters.
### `blend`
Blends the result of two other filters.
Keys:
```yaml
# No default, value for example purpose only.
filter1:
filter_type: simple
# No default, value for example purpose only.
filter2:
filter_type: simple
```
### `list`
A list of filters. The output of the previous is given to the next as input.
The `list` filter's blend applies to the output from the final filter in the list.
Keys:
```yaml
# Values for example only, default is an empty list.
filters:
- filter_type: simple
strength: 1.0
- filter_type: simple
strength: 1.0
```
### `bleh_enhance`
Available if you have the [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) node pack installed. See:
https://github.com/blepping/ComfyUI-bleh#enhancement-types
Keys:
```yaml
enhance_mode: null
enhance_scale: 1.0
```
### `bleh_ops`
Available if you have the [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) node pack installed. See:
https://github.com/blepping/ComfyUI-bleh#blehblockops
Keys:
```yaml
# May be specified as a string containing the YAML rule definitions or inline.
# Values for example only, default is an empty list of ops.
ops:
- if:
to_percent: 0.5
ops: # Not recommended to actually do this.
- [flip, { direction: h }]
- [roll, { direction: channels, amount: -2 }]
```
### `sonar_power_filter`
Available if you have [ComfyUI-sonar](https://github.com/blepping/ComfyUI-sonar) installed. See:
https://github.com/blepping/ComfyUI-sonar/blob/main/docs/advanced_power_noise.md
Keys:
```yaml
power_filter:
mix: 1.0
normalization_factor: 1.0
common_mode: 0.0
channel_correlation: "1,1,1,1,1,1"
alpha: 0.0
min_freq: 0.0
max_freq: 0.7071
stretch: 1.0
rotate: 0.0
pnorm: 2.0
scale: 1.0
compose_mode: max
# If specified should be another power filter definition.
compose_with: null
```
+18
View File
@@ -0,0 +1,18 @@
from . import types, expression, handler, util, validation
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",
)
+239
View File
@@ -0,0 +1,239 @@
import re
import operator
from .parser import Parser, ParserSpec, ParseError
from .types import (
Empty,
ExpBase,
ExpOp,
ExpBinOp,
ExpSym,
ExpStatements,
ExpFunAp,
ExpTuple,
ExpDict,
ExpKV,
)
class Expression:
EXPR_RE = re.compile(
r"""
\s*
(
\d+ # Possibly negative numeric literal
(?: \. \d* )? # Floating point
(?: e [+-] \d+)? # Scientific notation
| (?: \*\* | // ) # Doubled operators
| [<>]=? # Relative comparison
| [!=]= # Equality
| (?: \|\| | && ) # Logic
| [-+*/|!(),] # Operators
| :> # Key value binop
| ;
| \[ | ]
| \.\.\.
| '[\w.]+ # Symbol
| `?[a-z][\w.]*`? # Function/variable names
)
\s*
""",
re.I | re.S | re.X | re.A,
)
def __init__(self, toks):
if isinstance(toks, str):
toks = tuple(self.tokenize(toks))
self.expr = Parser(ExprParserSpec(), iter(toks)).go()
def __repr__(self):
return f"<Expr{self.expr!r}>"
def __call__(self, *args, **kwargs):
return self.eval(*args, **kwargs)
def eval(self, handlers, *args, **kwargs):
print("\nEVAL", self.expr)
if not isinstance(self.expr, ExpBase):
return self.expr
return self.expr.eval(handlers, *args, **kwargs)
def __len__(self):
return len(self.expr)
def pretty_string(self, depth=0):
sval = (
repr(self.expr)
if not isinstance(self.expr, ExpBase)
else self.expr.pretty_string(depth=depth + 1)
)
pad = " " * (depth + 1) * 2
return f"<Expr:\n{pad}{sval}\n{pad[:-2]}>"
FIXUP = {"true": True, "false": False, "...": Ellipsis, "none": None}
@classmethod
def fixup_token(cls, t):
if t == "":
return t
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):
yield from (cls.fixup_token(m.group(1)) for m in cls.EXPR_RE.finditer(s))
STATIC_OP_HANDLERS = {
"+": operator.add,
"-": operator.sub,
"*": operator.mul,
"/": operator.truediv,
"//": operator.floordiv,
"**": operator.pow,
"%": operator.mod,
"add": operator.add,
"sub": operator.sub,
"mul": operator.mul,
"div": operator.truediv,
"idiv": operator.floordiv,
"pow": operator.pow,
"mod": operator.mod,
}
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(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 = STATIC_OP_HANDLERS.get(op)
if h is None:
return ExpFunAp(op, args, kwargs)
return h(*args)
class ExprParserSpec(ParserSpec):
def __init__(self):
super().__init__()
self.populate()
@staticmethod
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)
})
@staticmethod
def null_constant(p, token, bp):
return token
@staticmethod
def null_paren(p, token, bp):
result = p.parse_until(bp) if p.token != ")" else ExpTuple()
if p.token == ",":
p.advance()
p.expect(")")
return result
@staticmethod
def null_prefixop(p, token, bp):
val = p.parse_until(bp)
return make_funap(token, ExpTuple((val,)))
@classmethod
def left_binop(cls, p, token, left, bp):
return make_funap(token, *cls.split_funap_args((left, p.parse_until(bp))))
@staticmethod
def left_kv(p, token, left, bp):
if not isinstance(left, (ExpOp, ExpSym)):
raise ParseError(f"{left!r} is not a valid key")
return ExpKV(left, p.parse_until(bp))
@classmethod
def left_funcall(cls, p, token, left, bp):
if not isinstance(left, ExpOp):
raise ParseError(f"{left!r} is not a valid function/variable name")
args = []
while p.lexer and p.token != ")":
args.append(p.parse_until(1))
if p.token == ",":
p.advance()
p.expect(")")
return make_funap(left, *cls.split_funap_args(args))
@staticmethod
def left_comma(p, token, left, bp):
if p.token == ")":
return left if isinstance(left, ExpTuple) else ExpTuple((left,))
r = p.parse_until(bp)
return ExpTuple((*left, r) if isinstance(left, ExpTuple) else (left, r))
@staticmethod
def left_semicolon(p, token, left, bp):
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)
if isinstance(left, ExpStatements)
else ExpTuple((left, r))
)
@staticmethod
def left_index(p, token, left, bp):
idx = p.parse_until(0)
p.expect("]")
return make_funap("index", ExpTuple((idx, left)))
@staticmethod
def get_type(token):
if isinstance(token, (int, float)):
return "number"
if isinstance(token, ExpSym):
return "sym"
if isinstance(token, ExpBinOp):
return "binop"
if isinstance(token, ExpOp) and token[0].isalpha():
return "op"
return token
def populate(self):
self.add_left(31, self.left_funcall, ("(",))
self.add_left(31, self.left_index, ("[",))
self.add_leftright(29, self.left_binop, ("**",))
self.add_null(27, self.null_prefixop, ("+", "-", "!"))
self.add_left(25, self.left_binop, ("*", "/"))
self.add_left(23, self.left_binop, ("+", "-"))
self.add_left(22, self.left_binop, ("binop",))
self.add_left(19, self.left_binop, ("<", ">", "<=", ">="))
self.add_left(19, self.left_binop, ("==", "!="))
self.add_left(9, self.left_binop, ("&&",))
self.add_left(7, self.left_binop, ("||",))
self.add_left(6, self.left_kv, (":>",))
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, (")", "]"))
+364
View File
@@ -0,0 +1,364 @@
import operator
from .validation import ValidateArg, Arg, ValidateError
from .types import Empty, ExpDict
from .util import torch
class HandlerError(Exception):
pass
class BaseHandler:
input_validators = ()
def __init__(self):
self.input_validators_by_key = {
v.name: (idx, v) for idx, v in enumerate(self.input_validators)
}
def __call__(self, obj, *, getter):
try:
val = self.handle(obj, getter)
return self.validate_output(obj, val)
except Exception as exc:
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)
if str_key:
argidx, validator = self.input_validators_by_key.get(key, (-1, None))
else:
argidx, validator = (
key,
(
self.input_validators[key]
if key < len(self.input_validators)
else None
),
)
default = (
default
if default is not Empty or validator is None
else getattr(validator, "default", Empty)
)
if argidx >= 0 and argidx < len(obj.args):
eff_key = argidx
str_eff_key = False
elif str_key:
eff_key = key
str_eff_key = True
else:
raise ValidateError(
f"Error validating input argument {key}, out of range for actual function arguments"
)
if getter is None:
if str_eff_key:
val = obj.kwargs.get(eff_key)
else:
val = default if eff_key > len(obj.args) else obj.args[eff_key]
else:
val = getter(eff_key, default=default)
if validator is None:
return val
try:
return validator(key, val)
except ValidateError as exc:
raise ValidateError(
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)
def safe_get_all(self, obj, getter=None, *, default=Empty):
return self.safe_get_multi(
(v.name for v in self.input_validators), obj, getter, default=default
)
def handle(self, obj, getter):
raise NotImplementedError
def validate_output(self, obj, value):
return value
class BinopLogicHandler(BaseHandler):
input_validators = (
Arg.present("lhs"),
Arg.present("rhs"),
)
def validate_output(self, obj, value):
return operator.truth(value)
class OrHandler(BinopLogicHandler):
def handle(self, obj, getter):
return operator.truth(
self.safe_get("lhs", obj, getter=getter)
) or operator.truth(self.safe_get("rhs", obj, getter=getter))
class AndHandler(BinopLogicHandler):
def handle(self, obj, getter):
return operator.truth(
self.safe_get("lhs", obj, getter=getter)
) and operator.truth(self.safe_get("rhs", obj, getter=getter))
class AllHandler(BinopLogicHandler):
input_validators = ()
def handle(self, obj, getter):
return all(
operator.truth(self.safe_get(idx, obj, getter=getter))
for idx in range(len(obj.args))
) and all(
operator.truth(self.safe_get(key, obj, getter=getter)) for key in obj.kwargs
)
class AnyHandler(BinopLogicHandler):
def handle(self, obj, getter):
return any(
operator.truth(self.safe_get(idx, obj, getter=getter))
for idx in range(len(obj.args))
) or any(
operator.truth(self.safe_get(key, obj, getter=getter)) for key in obj.kwargs
)
class EqHandler(BinopLogicHandler):
def handle(self, obj, getter):
a1, a2 = self.safe_get_all(obj, getter)
if isinstance(a1, torch.Tensor) and isinstance(a2, torch.Tensor):
return torch.equal(a1, a2)
return a1 == a2
class NeqHandler(BinopLogicHandler):
def handle(self, *args, **kwargs):
return not super().handle(*args, **kwargs)
class NotHandler(BinopLogicHandler):
input_validators = (Arg.present("value"),)
def handle(self, obj, getter):
return not operator.truth(self.safe_get("value", obj, getter=getter))
class IfHandler(BaseHandler):
input_validators = (
Arg.present("condition"),
Arg.present("then"),
Arg.present("else"),
)
def handle(self, obj, getter):
if operator.truth(self.safe_get("condition", obj, getter=getter)):
return self.safe_get("then", obj, getter=getter)
return self.safe_get("else", obj, getter=getter)
class BetweenHandler(BaseHandler): # Inclusive
input_validators = (
Arg.numeric("value"),
Arg.numeric("from", 0.0),
Arg.numeric("to"),
)
def handle(self, obj, getter):
value, low, high = self.safe_get_all(obj, getter)
return low <= value <= high
class SimpleMathHandler(BaseHandler):
input_validators = (Arg.numeric("lhs"), Arg.numeric("rhs"))
def __init__(self, handler):
super().__init__()
self.handler = handler
def validate_output(self, obj, value):
return ValidateArg.validate_numeric(-1, value)
def handle(self, obj, getter):
args = (
self.safe_get(idx, obj, getter=getter)
for idx in range(len(self.input_validators))
)
return self.handler(*args)
class MinusHandler(SimpleMathHandler):
input_validators = (Arg.numeric("lhs"), Arg.numeric("rhs", default=Empty))
__init__ = BaseHandler.__init__
def handle(self, obj, getter):
lhs, rhs = self.safe_get_all(obj, getter)
if rhs is Empty:
return operator.neg(lhs)
return operator.sub(lhs, rhs)
class RelComparisonHandler(SimpleMathHandler):
def validate_output(self, obj, value):
return operator.truth(value)
class UnarySimpleMathHandler(SimpleMathHandler):
input_validators = (Arg.numeric("lhs"),)
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.handlers
def validate_output(self, obj, value):
return operator.truth(value)
class GetHandler(BaseHandler):
input_validators = (
Arg.string("name"),
Arg.present("fallback"),
)
def handle(self, obj, getter):
key = self.safe_get("name", obj, getter=getter)
h = getter.handlers.get(key)
if h is None:
return self.safe_get("fallback", obj, getter=getter)
return h(getter.handlers, *getter.args, **getter.kwargs)
class S_Handler(BaseHandler):
input_validators = (
Arg.integer("start", None),
Arg.integer("end", None),
Arg.integer("step", None),
)
def handle(self, obj, getter):
return slice(*self.safe_get_all(obj, getter=getter))
class IndexHandler(BaseHandler):
input_validators = (
Arg.present("index"),
Arg.one_of(
"value", (ValidateArg.validate_sequence, ValidateArg.validate_tensor)
),
)
def handle(self, obj, getter):
idx, value = self.safe_get_all(obj, getter=getter)
return value[idx]
class MinHandler(BaseHandler):
input_validators = (Arg.numscalar_sequence("values"),)
def handle(self, obj, getter):
return min(*self.safe_get("values", obj, getter))
def validate_output(self, obj, value):
return ValidateArg.validate_numeric(-1, value)
class MaxHandler(MinHandler):
def handle(self, obj, getter):
return max(*self.safe_get("values", obj, getter))
class UnsafeCallHandler(BaseHandler):
input_validators = (Arg.present("__callable"),)
def handle(self, obj, getter):
if "__callable" in obj.kwargs:
raise ValueError(
"unsafe_call does not support passing the callable via keyword arg"
)
fun = self.safe_get("__callable", obj, getter)
if not callable(fun):
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)
class DictHandler(BaseHandler):
def handle(self, obj, getter):
if len(obj.args):
raise ValueError("Non-KV items passed to dict constructor")
return ExpDict({k: self.safe_get(k, obj, getter) for k in obj.kwargs.keys()})
class CommentHandler(BaseHandler):
def handle(self, obj, getter):
return None
LOGIC_HANDLERS = {
"||": OrHandler(),
"&&": AndHandler(),
"==": EqHandler(),
"!=": NeqHandler(),
"not": NotHandler(),
"if": IfHandler(),
"all": AllHandler(),
"any": AnyHandler(),
}
for k, alias in (
("||", "or"),
("&&", "and"),
("==", "eq"),
("!=", "neq"),
):
LOGIC_HANDLERS[alias] = LOGIC_HANDLERS[k]
MATH_HANDLERS = {
"+": 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),
"min": MinHandler(),
"max": MaxHandler(),
}
for k, alias in (
("+", "add"),
("-", "sub"),
("*", "mul"),
("/", "div"),
("//", "idiv"),
("**", "pow"),
):
MATH_HANDLERS[alias] = MATH_HANDLERS[k]
MISC_HANDLERS = {
"is_set": IsSetHandler(),
"get": GetHandler(),
"index": IndexHandler(),
"s_": S_Handler(),
"unsafe_call": UnsafeCallHandler(),
"dict": DictHandler(),
"comment": CommentHandler(),
}
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
+112
View File
@@ -0,0 +1,112 @@
class ParseError(Exception):
pass
# Pratt parsing referenced from https://github.com/andychu/pratt-parsing-demo
class ParserSpec:
@staticmethod
def null_error(p, token, bp):
raise ParseError(f"{token!r} cannot be used in prefix position")
@staticmethod
def left_error(p, token, bp):
raise ParseError(f"{token!r} cannot be used in infix position")
class LeftInfo:
def __init__(self, led=None, lbp=0, rbp=0):
self.led, self.lbp, self.rbp = led or ParserSpec.left_error, lbp, rbp
class NullInfo:
def __init__(self, nud=None, bp=0):
self.nud, self.bp = nud or ParserSpec.null_error, bp
def __init__(self):
self.null_lookup = {}
self.left_lookup = {}
def add_null(self, bp, nud, tokens):
for token in tokens:
self.null_lookup[token] = self.NullInfo(nud, bp)
if token not in self.left_lookup:
self.left_lookup[token] = self.LeftInfo()
def add_led(self, lbp, rbp, led, tokens):
for token in tokens:
self.left_lookup[token] = self.LeftInfo(led, lbp, rbp)
if token not in self.null_lookup:
self.null_lookup[token] = self.NullInfo(self.null_error)
def add_left(self, bp, led, tokens):
return self.add_led(bp, bp, led, tokens)
def add_leftright(self, bp, led, tokens):
return self.add_led(bp, bp - 1, led, tokens)
def lookup(self, token, is_left):
result = (self.left_lookup if is_left else self.null_lookup).get(token)
if result is None:
raise ParseError(f"Unexpected token {token!r}")
return result
@staticmethod
def get_type(token):
if isinstance(token, (int, float)):
return "number"
if isinstance(token, str) and token.isidentifier():
return "op"
return token
class Parser:
def __init__(self, spec, lexer):
self.spec = spec
self.lexer = lexer
self.token = None
self.token_type = None
self.pos = -1
def advance(self):
if self.lexer is None:
self.token_type = self.token = None
return None
try:
self.token = next(self.lexer)
self.token_type = self.spec.get_type(self.token)
self.pos += 1
except StopIteration:
self.token = self.token_type = self.lexer = None
return self.token
def expect(self, val):
if val is not None and (self.lexer is None or self.token != val):
raise ParseError(f"expected {val!r}, got {self.token!r}")
return self.advance()
def parse_until(self, rbp):
if self.lexer is None:
raise ParseError("unexpected end of input")
spec = self.spec
token, token_type = self.token, self.token_type
self.advance()
ni = spec.lookup(token_type, False)
node = ni.nud(self, token, ni.bp)
while self.lexer:
token, token_type = self.token, self.token_type
li = spec.lookup(token_type, True)
if rbp >= li.lbp:
break
self.advance()
node = li.led(self, token, node, li.rbp)
return node
def go(self):
self.advance()
try:
result = self.parse_until(0)
except ParseError as exc:
raise ParseError(
f"pos {self.pos} at token {self.token!r}: parse error: {exc}"
) from None
if self.lexer:
raise ParseError(f"pos {self.pos}: unexpected end of input")
return result
+220
View File
@@ -0,0 +1,220 @@
class Empty:
def __bool__(self):
return False
class ExpBase:
def __bool__(self):
return True
def pretty_string(self, *, depth=0):
return repr(self)
def eval(self, *args, **kwargs):
return self
def clone(self, *, mapper=None):
return self if not mapper else mapper(self)
class ExpOp(str, ExpBase):
__slots__ = ()
def eval(self, handlers, *args, **kwargs):
h = handlers.get(self)
if h is None:
raise KeyError(f"No handler for op/var {self}")
return h(handlers, *args, **kwargs)
class ExpBinOp(ExpOp):
__slots__ = ()
class ExpSym(str, ExpBase):
__slots__ = ()
def __repr__(self):
return f"'{self}"
class ExpTuple(tuple, ExpBase):
__slots__ = ()
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)
if isinstance(val, ExpBase):
return val.eval(handlers, *args, **kwargs)
return val
def pretty_string(self, depth=0):
vals = (
repr(v) if not isinstance(v, ExpBase) else v.pretty_string(depth=depth + 1)
for v in self
)
pad = " " * (depth + 1) * 2
nlpad = f",\n{pad}"
return f"(\n{pad}{nlpad.join(vals)}\n{pad[:-2]})"
def eval(self, handlers, *args, **kwargs):
return tuple(
v.eval(handlers, *args, **kwargs) if isinstance(v, ExpBase) else v
for v in self
)
class ExpKV(ExpBase):
__slots__ = ("k", "v")
def __init__(self, k, v):
self.k = k
self.v = v
class ExpDict(dict, ExpBase):
__slots__ = ()
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)
if isinstance(val, ExpBase):
return val.eval(handlers, *args, **kwargs)
return val
def pretty_string(self, depth=0):
vals = (
f"{k}: {v!r}"
if not isinstance(v, ExpBase)
else f"{k}: {v.pretty_string(depth=depth + 1)}"
for k, v in self.items()
)
pad = " " * (depth + 1) * 2
nlpad = f",\n{pad}"
return f"{{\n{pad}{nlpad.join(vals)}\n{pad[:-2]}}}"
def eval(self, handlers, *args, **kwargs):
return {
k: v.eval(handlers, *args, **kwargs) if isinstance(v, ExpBase) else v
for k, v in self.items()
}
popitem = pop
update = pop
clear = pop
__delitem__ = pop
__setitem__ = pop
__ior__ = pop
class ExpStatements(ExpBase):
def __init__(self, statements):
if not isinstance(statements, ExpTuple) or not len(statements):
raise ValueError("Must have at least one statement")
self.statements = statements
def eval(self, handlers, *args, **kwargs):
result = Empty
for stmt in self.statements:
result = (
stmt.eval(handlers, *args, **kwargs)
if isinstance(stmt, ExpBase)
else stmt
)
return result
def __repr__(self):
return f"@{self.statements}"
class ExprGetter:
GetterEmpty = Empty
# class GetterEmpty:
# def __bool__(self):
# return False
def __init__(self, obj, handlers, *args, **kwargs):
self.obj = obj
self.handlers = handlers
self.args = args
self.kwargs = kwargs
def __call__(self, k, *, default=GetterEmpty):
obj = self.obj
result = (
obj.kwargs.get_eval(
k, self.handlers, *self.args, default=default, **self.kwargs
)
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 ExpFunAp(ExpBase):
__slots__ = ("name", "args", "kwargs")
def __init__(self, name, args=None, kwargs=None):
self.name = name
self.args = args if args is not None else ExpTuple()
self.kwargs = kwargs if kwargs is not None else ExpDict()
def eval(self, handlers, *args, **kwargs):
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):
return self.__class__(self.name, self.args.clone(), self.kwargs.clone())
def pretty_string(self, depth=0):
pad = " " * (depth + 1) * 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):
return (
f"<FUNAP:{self.name}{self.args}{f", {self.kwargs}" if self.kwargs else ''}>"
)
class ExpBoundFunAp(ExpFunAp):
__slots__ = ("fun",)
def __init__(self, name, fun, args, kwargs):
super().__init__(name, args, kwargs)
self.fun = fun
def eval(self, handlers, *args, **kwargs):
def get_evaled(k, default=None):
return (
self.kwargs.get_eval(k, handlers, *args, default=default, **kwargs)
if isinstance(k, str)
else self.args.get_eval(k, handlers, *args, **kwargs)
)
return self.fun(self.name, self.args, *args, getter=get_evaled, **kwargs)
__all__ = (
"ExpBase",
"ExpOp",
"ExpBinOp",
"ExpSym",
"ExpTuple",
"ExpKV",
"ExpDict",
"ExpFunAp",
"ExpBoundFunAp",
)
+36
View File
@@ -0,0 +1,36 @@
import itertools
try:
import torch
except ImportError:
# To facilitate testing.
class torch:
class Tensor:
pass
class WrapGenerator:
def __init__(self, g):
self.g = g
self._value = None
self.ready = False
@property
def value(self):
if not self.ready:
raise ValueError("Value not ready")
return self._value
def __iter__(self):
self._value = yield from self.g
self.ready = True
return self._value
def split_iterable(seq, pred):
it = iter(seq)
while True:
toks = tuple(itertools.takewhile(pred, it))
if toks == ():
break
yield toks
+190
View File
@@ -0,0 +1,190 @@
import functools
from .util import torch
class Arg:
__slots__ = ("name", "default", "validator")
class Empty:
pass
def __init__(self, name, default=Empty, *, validator=None):
self.name = name
self.default = default
self.validator = validator
def __call__(self, _key, value, *args, **kwargs):
return self.validate(value, *args, **kwargs)
def validate(self, value):
# 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:
return self.validator(self.name, value) if self.validator else value
except ValidateError as exc:
raise ValidateError(f"Failed to validate argument {self.name}: {exc}")
@classmethod
def tensor(cls, name):
return cls(name, validator=ValidateArg.validate_tensor)
@classmethod
def numeric(cls, name, default=Empty):
return cls(name, default=default, validator=ValidateArg.validate_numeric)
@classmethod
def numeric_scalar(cls, name, default=Empty):
return cls(name, default=default, validator=ValidateArg.validate_numeric_scalar)
@classmethod
def integer(cls, name, default=Empty):
return cls(name, default=default, validator=ValidateArg.validate_integer)
@classmethod
def numscalar_sequence(cls, name, default=Empty):
return cls(
name, default=default, validator=ValidateArg.validate_numscalar_sequence
)
@classmethod
def sequence(cls, name, default=Empty, *, item_validator=None):
return cls(
name,
default=default,
validator=functools.partial(
ValidateArg.validate_sequence, item_validator=item_validator
),
)
@classmethod
def string(cls, name, default=Empty):
return cls(name, default=default, validator=ValidateArg.validate_string)
@classmethod
def boolean(cls, name, default=Empty):
return cls(name, default=default, validator=ValidateArg.validate_boolean)
@classmethod
def present(cls, name):
return cls(name, validator=ValidateArg.validate_passthrough)
@classmethod
def one_of(cls, name, validators, *, default=Empty):
def validate(idx, val):
for validator in validators:
try:
return validator(idx, val)
except ValidateError:
continue
raise ValidateError(
f"Failed to validate argument at {idx} of type {type(val)}"
)
return cls(name, default=default, validator=validate)
class ValidateError(Exception):
pass
class ValidateArg:
__slots__ = ("valfuns", "groupfun", "kwargs", "kwargslist")
def __init__(self, name, *args, kwargslist=(), group=all, **kwargs):
if not isinstance(name, (list, tuple)):
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")
self.groupfun = group
self.kwargs = kwargs
self.kwargslist = kwargslist if kwargslist is not None else {}
def __call__(self, *args, **kwargs):
kalen = len(self.kwargslist)
return self.groupfun(
vf(
*args,
**(self.kwargslist if idx < kalen else {}),
**self.kwargs,
)
for idx, vf in enumerate(self.valfuns)
)
@staticmethod
def validate_numeric(idx, val):
if not isinstance(val, (int, float, torch.Tensor)):
raise ValidateError(
f"Expected numeric or tensor argument at {idx}, got {type(val)}"
)
return val
@classmethod
def validate_numeric_scalar(cls, idx, val):
if not isinstance(val, (int, float)):
raise ValidateError(f"Expected numeric argument at {idx}, got {type(val)}")
return val
@classmethod
def validate_integer(cls, idx, val):
if not isinstance(val, int):
raise ValidateError(f"Expected integer argument at {idx}, got {type(val)}")
return val
@staticmethod
def validate_tensor(idx, val):
if not isinstance(val, torch.Tensor):
raise ValidateError(f"Expected tensor argument at {idx}, got {type(val)}")
return val
@staticmethod
def validate_sequence(idx, val, *, item_validator=None):
if not isinstance(val, (list, tuple)):
raise ValidateError(f"Expected sequence argument at {idx}, got {type(val)}")
if item_validator is None:
return val
try:
return tuple(item_validator(iidx, v) for iidx, v in enumerate(val))
except ValidateError as exc:
raise ValidateError(f"Item validation failed for in sequence: {exc}")
@classmethod
def validate_numscalar_sequence(cls, idx, val):
return cls.validate_sequence(
idx, val, item_validator=cls.validate_numeric_scalar
)
# @classmethod
# def validate_numscalar_sequence(cls, idx, val):
# if not isinstance(val, (list, tuple)):
# raise ValidateError(f"Expected sequence argument at {idx}, got {type(val)}")
# try:
# _ = all(
# cls.validate_numeric_scalar(f"{idx}[{i}]", v) is not None
# for i, v in enumerate(val)
# )
# except ValidateError as exc:
# raise ValidateError(
# f"Expected numeric sequence argument at {idx}, got {type(val)}: {exc}"
# )
# return val
@classmethod
def validate_string(cls, idx, val):
if not isinstance(val, str):
raise ValidateError(f"Expected string 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_passthrough(cls, idx, val):
return val
+623
View File
@@ -0,0 +1,623 @@
import os
import torch
import numpy as np
from . import expression as expr
from . import latent
from .external import MODULES as EXT
from .utils import scale_noise, resolve_value
ALLOW_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_UNSAFE_EXPRESSIONS") is not None
ALLOW_ALL_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_ALL_UNSAFE") is not None
EXT_BLEH = EXT.get("bleh")
EXT_SONAR = EXT.get("sonar")
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,
}
HANDLERS = {}
class NormHandler(expr.BaseHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.numeric("factor", 1.0),
expr.Arg.numscalar_sequence("dim", (-3, -2, -1)),
)
def handle(self, obj, getter):
tensor, factor, dim = self.safe_get_all(obj, getter)
return scale_noise(tensor, factor, normalize_dims=dim)
validate_output = expr.Arg.tensor("output")
class MeanHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.numscalar_sequence("dim", (-3, -2, -1)),
)
def handle(self, obj, getter):
tensor, dim = self.safe_get_all(obj, getter)
return tensor.mean(keepdim=True, dim=dim)
class StdHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.numscalar_sequence("dim", (-3, -2, -1)),
)
def handle(self, obj, getter):
tensor, dim = self.safe_get_all(obj, getter)
return tensor.std(keepdim=True, dim=dim)
class RollHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.numeric_scalar("amount", 0.5),
expr.Arg.numscalar_sequence("dim", (-2,)),
)
def handle(self, obj, getter):
tensor, amount, dim = self.safe_get_all(obj, getter)
if isinstance(amount, float) and amount < 1.0 and amount > -1.0:
if len(dim) > 1:
raise ValueError(
"Cannot use percentage based amount with multiple roll dimensions",
)
amount = int(tensor.shape[dim[0]] * amount)
return tensor.roll(amount, dims=dim)
class FlipHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.integer("dim"),
expr.Arg.boolean("mirror", False),
)
def handle(self, obj, getter):
tensor, dim, mirror = self.safe_get_all(obj, getter)
if dim < 0:
dim += tensor.ndim
if dim < 0 or dim >= tensor.ndim:
raise ValueError(
f"Dimension out of range, wanted {dim}, tensor has {tensor.ndim} dimension(s)"
)
if not mirror:
return torch.flip(tensor, (dim,))
result = tensor.detach().clone()
pivot = tensor.shape[dim] // 2
out_slice = (
np.s_[:] if d != dim else np.s_[pivot:] for d in range(tensor.ndim)
)
in_slice = (np.s_[:] if d != dim else np.s_[:pivot] for d in range(tensor.ndim))
result[*out_slice] = torch.flip(tensor[*in_slice], dims=(dim,))
return result
class BlendHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor1"),
expr.Arg.tensor("tensor2"),
expr.Arg.numeric("scale", 0.5),
expr.Arg.string("mode", "lerp"),
)
def handle(self, obj, getter):
t1, t2, scale, mode = self.safe_get_all(obj, getter)
blend_handler = BLENDING_MODES.get(mode)
if not blend_handler:
raise KeyError(f"Unknown blend mode {mode!r}")
return blend_handler(t1, t2, scale)
class ContrastAdaptiveSharpeningHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.numeric("scale", 0.5),
)
def handle(self, obj, getter):
t, scale = self.safe_get_all(obj, getter)
return latent.contrast_adaptive_sharpening(t, scale)
class ScaleHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.one_of(
"scale",
(
expr.ValidateArg.validate_numeric_scalar,
expr.ValidateArg.validate_numscalar_sequence,
),
),
expr.Arg.string("mode", "bicubic"),
expr.Arg.boolean("absolute_scale", False),
)
def handle(self, obj, getter):
t, scale, mode, abs_scale = self.safe_get_all(obj, getter)
if isinstance(scale, (list, tuple)):
if len(scale) != 2:
raise ValueError(
"When passing scale as a tuple, it must be in the form (h, w)"
)
else:
scale = (scale, scale)
if abs_scale:
scale = tuple(int(v) for v in scale)
else:
scale = (int(t.shape[-1] * scale[0]), int(t.shape[-2] * scale[1]))
if not all(v > 0 for v in scale):
raise ValueError(f"Invalid scale: scale values must be > 0, got: {scale!r}")
return latent.scale_samples(t, scale[1], scale[0], mode=mode)
class NoiseHandler(NormHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.string("type", "gaussian"),
)
def handle(self, obj, getter):
t, typ = self.safe_get_all(obj, getter)
ctx = getter.handlers
smin, smax, s, sn = (
h(ctx, *getter.args, **getter.kwargs) if h is not None else None
for h in (
ctx.get(k) for k in ("sigma_min", "sigma_max", "sigma", "sigma_next")
)
)
ns = latent.get_noise_sampler(typ, t, smin, smax, normalized=False)
return ns(s, sn)
class UnsafeTorchTensorMethodHandler(NormHandler):
input_validators = (
expr.Arg.tensor("__tensor"),
expr.Arg.string("__method"),
)
if ALLOW_ALL_UNSAFE:
class AlwaysContains:
def __contains__(self, k):
return True
whitelist = AlwaysContains()
elif ALLOW_UNSAFE:
whitelist = {
"abs",
"absolute",
"acos",
"acosh",
"add",
"addbmm",
"addcdiv",
"addcmul",
"addmm",
"addmv",
"addr",
"adjoint",
"all",
"allclose",
"amax",
"amin",
"aminmax",
"angle",
"any",
"arccos",
"arccosh",
"arcsin",
"arcsinh",
"arctan",
"arctan2",
"arctanh",
"argmax",
"argmin",
"argsort",
"argwhere",
"as_strided",
"asin",
"asinh",
"atan",
"atan2",
"atanh",
"baddbmm",
"bernoulli",
"bincount",
"bitwise_and",
"bitwise_left_shift",
"bitwise_not",
"bitwise_or",
"bitwise_right_shift",
"bitwise_xor",
"bmm",
"broadcast_to",
"ceil",
"cholesky",
"cholesky_inverse",
"cholesky_solve",
"chunk",
"clamp",
"clip",
"clone",
"conj",
"conj_physical",
"contiguous",
"copysign",
"corrcoef",
"cos",
"cosh",
"count_nonzero",
"cov",
"cross",
"cummax",
"cummin",
"cumprod",
"cumsum",
"deg2rad",
"det",
"detach",
"diag",
"diag_embed",
"diagflat",
"diagonal",
"diagonal_scatter",
"diff",
"digamma",
"dim",
"dist",
"div",
"divide",
"dot",
"dsplit",
"eq",
"equal",
"erf",
"erfc",
"erfinv",
"exp",
"expand",
"expand_as",
"expm1",
"fix",
"flatten",
"flip",
"fliplr",
"flipud",
"float_power",
"floor",
"floor_divide",
"fmax",
"fmin",
"fmod",
"frac",
"frexp",
"gather",
"gcd",
"ge",
"geqrf",
"ger",
"greater",
"greater_equal",
"gt",
"hardshrink",
"heaviside",
"histc",
"hsplit",
"hypot",
"i0",
"igamma",
"igammac",
"index_add",
"index_copy",
"index_fill",
"index_put",
"index_reduce",
"index_select",
"inner",
"inverse",
"isclose",
"isfinite",
"isinf",
"isnan",
"isneginf",
"isposinf",
"kthvalue",
"lcm()",
"ldexp",
"le",
"lerp",
"less",
"less_equal",
"lgamma",
"log",
"log10",
"log1p",
"log2",
"logaddexp",
"logaddexp2",
"logcumsumexp",
"logdet",
"logical_and",
"logical_not",
"logical_or",
"logical_xor",
"logit",
"logsumexp",
"lt",
"lu",
"lu_solve",
"masked_fill",
"masked_scatter",
"masked_select",
"matmul",
"matrix_exp",
"max",
"maximum",
"mean",
"median",
"min",
"minimum",
"mm",
"mode",
"moveaxis",
"movedim",
"msort",
"mul",
"multinomial",
"multiply",
"mv",
"mvlgamma",
"nan_to_num",
"nanmean",
"nanmedian",
"nanquantile",
"nansum",
"narrow",
"narrow_copy",
"ne",
"neg",
"negative",
"new_empty",
"new_full",
"new_ones",
"new_zeros",
"nextafter",
"nonzero",
"norm",
"not_equal",
"numel",
"orgqr",
"ormqr",
"outer",
"permute",
"polygamma",
"positive",
"pow",
"prod",
"qr",
"quantile",
"rad2deg",
"ravel",
"reciprocal",
"remainder",
"renorm",
"repeat",
"repeat_interleave",
"reshape",
"reshape_as",
"resolve_conj",
"resolve_neg",
"roll",
"rot90",
"round",
"rsqrt",
"scatter",
"scatter_add",
"scatter_reduce",
"select",
"select_scatter",
"sgn",
"sigmoid",
"sign",
"signbit",
"sin",
"sinc",
"sinh",
"slice_scatter",
"slogdet",
"smm",
"softmax",
"sort",
"sparse_mask",
"split",
"sqrt",
"square",
"squeeze",
"sspaddmm",
"std",
"stft",
"sub",
"subtract",
"sum",
"sum_to_size",
"svd",
"swapaxes",
"swapdims",
"t",
"take",
"take_along_dim",
"tan",
"tanh",
"tensor_split",
"tile",
"topk",
"transpose",
"triangular_solve",
"tril",
"triu",
"true_divide",
"trunc",
"unflatten",
"unfold",
"unique",
"unique_consecutive",
"unsqueeze",
"var",
"vdot",
"view",
"view_as",
"vsplit",
"where",
"xlogy",
}
else:
whitelist = set()
def handle(self, obj, getter):
if "__method" in obj.kwargs or "__tensor" in obj.kwargs:
raise ValueError(
"Tensor method call doesn't support passing method or tensor with keyword args"
)
tensor = self.safe_get("__tensor", obj, getter=getter)
method = self.safe_get("__method", obj, getter=getter)
args = (
self.safe_get(idx, obj, getter=getter) for idx in range(2, len(obj.args))
)
kwargs = {k: self.safe_get(k, obj, getter=getter) for k in obj.kwargs.keys()}
if method not in self.whitelist:
raise ValueError(f"Method {method} not whitelisted: cannot call")
methodfun = getattr(tensor, method, None)
if methodfun is None:
raise KeyError(f"Unknown method {method} for Torch tensor")
return methodfun(*args, **kwargs)
class UnsafeTorchHandler(expr.BaseHandler):
input_validators = (expr.Arg.string("path"),)
if not ALLOW_ALL_UNSAFE:
def handle(self, obj, getter):
raise ValueError("Unsafe Torch access not allowed")
else:
def handle(self, obj, getter):
path = self.safe_get("path", obj, getter)
keys = path.split(".")
if not keys or not all(k for k in keys):
raise ValueError(f"Bad path {path}")
return resolve_value(keys, torch)
if EXT_BLEH:
class BlehEnhanceHandler(expr.BaseHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.string("mode"),
expr.Arg.numeric_scalar("scale", 1.0),
)
output_validator = expr.Arg.tensor("output")
def handle(self, obj, getter):
tensor, mode, scale = self.safe_get_all(obj, getter)
return EXT_BLEH.latent_utils.enhance_tensor(
tensor, mode, scale=scale, adjust_scale=False
)
HANDLERS["t_bleh_enhance"] = BlehEnhanceHandler()
if EXT_SONAR:
class SonarPowerFilterHandler(expr.BaseHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.present("filter"),
)
output_validator = expr.Arg.tensor("output")
default_power_filter = {
"mix": 1.0,
"normalization_factor": 1.0,
"common_mode": 0.0,
"channel_correlation": "1,1,1,1,1,1",
}
@classmethod
def make_power_filter(cls, fdict, *, toplevel=True):
fdict = fdict.copy()
compose_with = fdict.pop("compose_with", None)
if compose_with:
if not isinstance(compose_with, dict):
raise TypeError("compose_with must be a dictionary")
fdict["compose_with"] = cls.make_power_filter(
compose_with, toplevel=False
)
topargs = {
k: fdict.pop(k, dv) for k, dv in cls.default_power_filter.items()
}
power_filter = EXT_SONAR.powernoise.PowerFilter(**fdict)
if not toplevel:
return power_filter
cc = topargs.get("channel_correlation")
if cc is not None:
if not isinstance(cc, (list, tuple)) or not all(
isinstance(v, (int, float)) for v in cc
):
raise TypeError(
"Bad channel correlation type: must be comma separated string or numeric sequence"
)
topargs["channel_correlation"] = ",".join(repr(v) for v in cc)
return EXT_SONAR.powernoise.PowerNoiseItem(
1, power_filter=power_filter, time_brownian=True, **topargs
)
def handle(self, obj, getter):
tensor, filter_def = self.safe_get_all(obj, getter)
if not isinstance(filter_def, dict):
raise TypeError("filter argument must be a dictionary")
power_filter = self.make_power_filter(filter_def)
filter_rfft = power_filter.make_filter(tensor.shape).to(
tensor.device, non_blocking=True
)
ns = power_filter.make_noise_sampler_internal(
tensor,
lambda *_unused, latent=tensor: latent,
filter_rfft,
normalized=False,
)
return ns(None, None)
HANDLERS["t_sonar_power_filter"] = SonarPowerFilterHandler()
TENSOR_OP_HANDLERS = {
"t_norm": NormHandler(),
"t_mean": MeanHandler(),
"t_std": StdHandler(),
"t_blend": BlendHandler(),
"t_roll": RollHandler(),
"t_flip": FlipHandler(),
"t_contrast_adaptive_sharpening": ContrastAdaptiveSharpeningHandler(),
"t_scale": ScaleHandler(),
"t_noise": NoiseHandler(),
"unsafe_tensor_method": UnsafeTorchTensorMethodHandler(),
"unsafe_torch": UnsafeTorchHandler(),
}
HANDLERS |= TENSOR_OP_HANDLERS
+18
View File
@@ -0,0 +1,18 @@
import contextlib
import importlib
MODULES = {}
with contextlib.suppress(ImportError, NotImplementedError):
bleh = importlib.import_module("custom_nodes.ComfyUI-bleh")
bleh_version = getattr(bleh, "BLEH_VERSION", -1)
if bleh_version < 1:
raise NotImplementedError
MODULES["bleh"] = bleh.py
with contextlib.suppress(ImportError, NotImplementedError):
sonar = importlib.import_module("custom_nodes.ComfyUI-sonar")
MODULES["sonar"] = sonar.py
__all__ = ("MODULES",)
+536
View File
@@ -0,0 +1,536 @@
import collections
import torch
from . import expression as expr
from . import expression_handlers
from .external import MODULES as EXT
from .utils import fallback
OD = collections.OrderedDict
EXT_BLEH = EXT.get("bleh")
EXT_SONAR = EXT.get("sonar")
if "bleh" in EXT:
BLENDING_MODES = EXT_BLEH.latent_utils.BLENDING_MODES
else:
BLENDING_MODES = {
"lerp": lambda a, b, t: (1 - t) * a + t * b,
}
BLENDING_MODES = BLENDING_MODES | {
"a_only": lambda a, b, t: a * t,
"b_only": lambda a, b, t: b * t,
}
FILTER = {}
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)
FILTER_HANDLERS = FilterHandlerCollection(
expr.BASIC_HANDLERS | expression_handlers.HANDLERS, {}
)
class FilterRefs:
def __init__(self, kvs=None):
self.kvs = fallback(kvs, {})
def get(self, k, default=None):
return self.kvs.get(k, default)
def __getitem__(self, k):
return self.kvs[k]
def __setitem__(self, k, v):
self.kvs[k] = v
def clone(self):
return self.__class__(self.kvs.copy())
def __or__(self, other):
return self.__class__(self.kvs | other.kvs)
def __ior__(self, other):
self.kvs |= other.kvs
return self
def __delitem__(self, k):
del self.kvs[k]
def __contains__(self, k):
return k in self.kvs
def __missing__(self, k):
return self.kvs.__missing__(k)
def __len__(self):
return len(self.kvs)
def __iter__(self):
return self.kvs.__iter__()
@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_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
if not have_current and len(ss.hist) > 0:
hist_offs = -1
elif len(ss.hist) > 1:
hist_offs = -2
else:
hist_offs = None
if hist_offs is not None:
hprev = ss.hist[hist_offs]
fr.kvs |= {f"{k}_prev": v for k, v in cls.from_mr(hprev).kvs.items()}
fr["d_prev"] = hprev.d
return fr
@classmethod
def from_mr(cls, mr):
return cls({
k: getattr(mr, ak)
for k, ak in (
("cond", "denoised_cond"),
("denoised", "denoised"),
("model_call", "call_idx"),
("sigma", "sigma"),
("uncond", "denoised_uncond"),
("x", "x"),
)
if getattr(mr, ak, None) is not None
})
@classmethod
def from_sr(cls, sr):
return cls({
k: getattr(sr, ak)
for k, ak in (
("cond", "denoised_cond"),
("denoised", "denoised"),
("noise", "noise_pred"),
("sigma_down", "sigma_down"),
("sigma_next", "sigma_next"),
("sigma_up", "sigma_up"),
("sigma", "sigma"),
("step", "step"),
("substep", "substep"),
("uncond", "denoised_uncond"),
("x", "x"),
)
if getattr(sr, ak, None) is not None
})
class Filter:
name = "unknown"
uses_ref = False
default_options = {
"enabled": True,
"when": None,
"input": "default",
"output": "default",
"ref": "default",
"final": "default",
"blend_mode": "lerp",
"strength": 1.0,
}
def __init__(self, **options):
self.options = options
self.set_options(self.default_options)
if self.when is not None:
self.when = expr.Expression(self.when)
for key in ("input", "output", "ref", "final"):
if not self.uses_ref and key == "ref":
continue
val = getattr(self, key)
val = make_filter(val) if isinstance(val, dict) else expr.Expression(val)
setattr(self, key, val)
if self.blend_mode not in BLENDING_MODES:
raise ValueError("Bad blend mode")
def set_options(self, defaults):
for k, v in defaults.items():
setattr(self, k, self.options.pop(k, v))
def apply(self, input_latent, default_ref=None, refs=None, **kwargs):
if not self.check_applies(refs):
return input_latent
refs = fallback(refs, FilterRefs()).clone()
latent = self.get_ref("input", input_latent, self.input, refs=refs)
refs["input"] = latent
if not self.uses_ref:
ref_latent = None
else:
ref_latent = (
self.get_ref("ref", default_ref, self.ref, refs=refs)
if default_ref is not None
else None
)
refs["ref"] = ref_latent
output_latent = self.get_ref(
"output",
self.filter(latent, ref_latent, refs=refs, **kwargs),
self.output,
refs=refs,
)
refs["output"] = output_latent
return self.get_ref(
"final",
BLENDING_MODES[self.blend_mode](
input_latent[: output_latent.shape[0]], output_latent, self.strength
),
self.final,
refs=refs,
)
def filter(self, latent, ref_latent, *, refs, **kwargs):
raise NotImplementedError
def check_applies(self, refs=None):
if not self.enabled:
return False
if self.when is None:
return True
refs = fallback(refs, FilterRefs())
matched = self.when.eval(FILTER_HANDLERS.clone_with_refs(refs))
# if matched:
# print("\nMATCH", self.name)
return matched
def get_ref(self, name, default_ref, ops, *, refs=None):
if isinstance(ops, 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_with_refs(refs))
class SimpleFilter(Filter):
name = "simple"
def filter(self, latent, *args, **kwargs):
return latent
class BlendFilter(Filter):
name = "blend"
default_options = Filter.default_options | {"filter1": None, "filter2": None}
def __init__(self, **kwargs):
super().__init__(**kwargs)
if not (isinstance(self.filter1, dict) and isinstance(self.filter2, dict)):
raise ValueError("Must set filter1 and filter2")
self.filter1 = make_filter(self.filter1)
self.filter2 = make_filter(self.filter2)
def filter(self, latent, ref_latent, *, refs, **kwargs):
if self.blend_mode == "lerp":
if self.strength == 0:
return self.filter1(latent, ref_latent, refs=refs, **kwargs)
if self.strength == 1:
return self.filter2(latent, ref_latent, refs=refs, **kwargs)
return BLENDING_MODES[self.blend_mode](
self.filter1.apply(latent, ref_latent, refs=refs, **kwargs),
self.filter2.apply(latent, ref_latent, refs=refs, **kwargs),
self.strength,
)
class ListFilter(Filter):
name = "list"
default_options = Filter.default_options | {"filters": ()}
def __init__(self, **kwargs):
super().__init__(**kwargs)
if not isinstance(self.filters, (list, tuple)):
raise ValueError("filters key must be a sequence")
self.filters = tuple(make_filter(filt) for filt in self.filters)
def filter(self, latent, ref_latent, *, refs, **kwargs):
if not self.filters:
return latent
for filt in self.filters:
latent = filt.apply(latent, ref_latent, refs=refs, **kwargs)
return latent
class NormalizeFilter(Filter):
name = "normalize"
uses_ref = True
default_options = Filter.default_options | {
"adjust_target": 0,
"balance_scale": 1.0,
"adjust_scale": 1.0,
"dims": (-2, -1),
}
def __init__(
self,
start_step=0,
end_step=9999,
phase="after",
adjust_target=0,
balance_scale=1.0,
adjust_scale=1.0,
dims=(-2, -1),
):
self.start_step = start_step
self.end_step = end_step
self.phase = phase.lower().strip() # before, after, all
if isinstance(adjust_target, str):
adjust_target = adjust_target.lower().strip()
if adjust_target not in ("x",):
raise ValueError("Bad target mean")
# "x", scalar or array matching mean dims
self.adjust_target = adjust_target
# multiplier on adjustment, scalar or array matching mean dims
self.adjust_scale = adjust_scale
self.balance_scale = balance_scale
self.dims = dims
def __call__(self, ss, sigma, latent, phase, orig_x=None):
if ss.step < self.start_step or ss.step > self.end_step:
return latent
if self.phase != "all" and phase != self.phase:
return latent
if self.adjust_target == "x" and orig_x is None:
raise ValueError("Can only use source x in after phase")
adjust_scale, balance_scale = (
torch.tensor(v, dtype=latent.dtype).to(latent)
if isinstance(v, (list, tuple))
else v
for v in (self.adjust_scale, self.balance_scale)
)
latent_mean = latent.mean(dim=self.dims, keepdim=True)
# print("MEAN", latent_mean)
latent = latent - latent_mean * balance_scale
if self.adjust_target == "x":
latent += orig_x.mean(dim=self.dims, keepdim=True) * adjust_scale
elif isinstance(self.adjust_target, (list, tuple)):
adjust_target = torch.tensor(self.adjust_target, dtype=latent.dtype).to(
latent
)
latent += adjust_target * adjust_scale
else:
latent += self.adjust_target * adjust_scale
return latent
class NormalizeFilter_:
def __init__(
self,
start_step=0,
end_step=9999,
phase="after",
adjust_target=0,
balance_scale=1.0,
adjust_scale=1.0,
dims=(-2, -1),
):
self.start_step = start_step
self.end_step = end_step
self.phase = phase.lower().strip() # before, after, all
if isinstance(adjust_target, str):
adjust_target = adjust_target.lower().strip()
if adjust_target not in ("x",):
raise ValueError("Bad target mean")
# "x", scalar or array matching mean dims
self.adjust_target = adjust_target
# multiplier on adjustment, scalar or array matching mean dims
self.adjust_scale = adjust_scale
self.balance_scale = balance_scale
self.dims = dims
def __call__(self, ss, sigma, latent, phase, orig_x=None):
if ss.step < self.start_step or ss.step > self.end_step:
return latent
if self.phase != "all" and phase != self.phase:
return latent
if self.adjust_target == "x" and orig_x is None:
raise ValueError("Can only use source x in after phase")
adjust_scale, balance_scale = (
torch.tensor(v, dtype=latent.dtype).to(latent)
if isinstance(v, (list, tuple))
else v
for v in (self.adjust_scale, self.balance_scale)
)
latent_mean = latent.mean(dim=self.dims, keepdim=True)
# print("MEAN", latent_mean)
latent = latent - latent_mean * balance_scale
if self.adjust_target == "x":
latent += orig_x.mean(dim=self.dims, keepdim=True) * adjust_scale
elif isinstance(self.adjust_target, (list, tuple)):
adjust_target = torch.tensor(self.adjust_target, dtype=latent.dtype).to(
latent
)
latent += adjust_target * adjust_scale
else:
latent += self.adjust_target * adjust_scale
return latent
Normalize = NormalizeFilter
if EXT_BLEH:
class BlehEnhanceFilter(Filter):
name = "bleh_enhance"
default_options = Filter.default_options | {
"enhance_mode": None,
"enhance_scale": 1.0,
}
def filter(self, latent, *args, **kwargs):
if self.enhance_mode is None or self.enhance_scale == 1:
return latent
return EXT_BLEH.latent_utils.enhance_tensor(
latent, self.enhance_mode, scale=self.enhance_scale, adjust_scale=False
)
class BlehOpsFilter(Filter):
name = "bleh_ops"
default_options = Filter.default_options | {"ops": ()}
def __init__(self, **kwargs):
super().__init__(**kwargs)
if isinstance(self.ops, (tuple, list)):
self.ops = EXT_BLEH.nodes.ops.RuleGroup(
tuple(
r
for rs in self.ops
for r in EXT_BLEH.nodes.ops.Rule.from_dict(rs)
)
)
return
if not isinstance(self.ops, str):
raise ValueError("ops key must be a YAML string or list of object")
self.ops = EXT_BLEH.nodes.ops.RuleGroup.from_yaml(self.ops)
def filter(self, latent, ref_latent, *args, refs=None, **kwargs):
if not self.ops:
return latent
refs = fallback(refs, {})
bops = EXT_BLEH.nodes.ops
state = {
bops.CondType.TYPE: bops.PatchType.LATENT,
bops.CondType.PERCENT: 0.0,
bops.CondType.BLOCK: -1,
bops.CondType.STAGE: -1,
bops.CondType.STEP: refs.get("step", 0),
bops.CondType.STEP_EXACT: refs.get("step", -1),
"h": latent,
"hsp": ref_latent,
"target": "h",
}
self.ops.eval(state, toplevel=True)
return state["h"]
FILTER |= {
"bleh_enhance": BlehEnhanceFilter,
"bleh_ops": BlehOpsFilter,
}
if EXT_SONAR:
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):
if not isinstance(args, dict):
raise TypeError(f"Bad type for filter: {type(args)}")
args = args.copy()
filter_type = args.pop("filter_type", "simple")
if not isinstance(filter_type, str):
raise ValueError("Missing or invalid filter_type")
filter_fun = FILTER.get(filter_type)
if filter_fun is None:
raise ValueError(f"Unknown filter_type: {filter_type}")
return filter_fun(**args)
FILTER |= {
"simple": SimpleFilter,
"blend": BlendFilter,
"list": ListFilter,
"normalize": NormalizeFilter,
}
+114
View File
@@ -0,0 +1,114 @@
import torch
import torch.nn.functional as F
from comfy.utils import bislerp
from .external import MODULES as EXT
# The following is modified to work with latent images of ~0 mean from https://github.com/Jamy-L/Pytorch-Contrast-Adaptive-Sharpening/tree/main.
def contrast_adaptive_sharpening(x, amount=0.8, *, epsilon=1e-06):
"""
Performs a contrast adaptive sharpening on the batch of images x.
The algorithm is directly implemented from FidelityFX's source code,
that can be found here
https://github.com/GPUOpen-Effects/FidelityFX-CAS/blob/master/ffx-cas/ffx_cas.h
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]
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
# padding = same by default
# Extracting the 3x3 neighborhood around each pixel
# a b c
# d e f
# g h i
a = x_padded[..., :-2, :-2]
b = x_padded[..., :-2, 1:-1]
c = x_padded[..., :-2, 2:]
d = x_padded[..., 1:-1, :-2]
e = x_padded[..., 1:-1, 1:-1]
f = x_padded[..., 1:-1, 2:]
g = x_padded[..., 2:, :-2]
h = x_padded[..., 2:, 1:-1]
i = x_padded[..., 2:, 2:]
# Computing contrast
cross = (b, d, e, f, h)
mn = on_abs_stacked(cross, torch.min, axis=0)
mx = on_abs_stacked(cross, torch.max, axis=0)
diag = (a, c, g, i)
mn2 = on_abs_stacked(diag, torch.min, axis=0)
mx2 = on_abs_stacked(diag, torch.max, axis=0)
mx = mx + mx2
mn = mn + mn2
# Computing local weight
inv_mx = torch.reciprocal(mx + epsilon) # 1/mx
amp = inv_mx * mn
# scaling
amp = torch.sqrt(amp)
w = -amp * (amount * (1 / 5 - 1 / 8) + 1 / 8)
# w scales from 0 when amp=0 to K for amp=1
# K scales from -1/5 when amount=1 to -1/8 for amount=0
# The local conv filter is
# 0 w 0
# w 1 w
# 0 w 0
div = torch.reciprocal(1 + 4 * w)
output = ((b + d + f + h) * w + e) * div
return output.real.clamp(x.min(), x.max())
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)
if "sonar" in EXT:
get_noise_sampler = EXT["sonar"].noise.get_noise_sampler
else:
def get_noise_sampler(noise_type, x, *_args: list, **_kwargs: dict):
if noise_type != "gaussian":
raise ValueError("Only gaussian noise supported")
return lambda _s, _sn: torch.randn_like(x)
+264
View File
@@ -0,0 +1,264 @@
from collections import namedtuple
import torch
import comfy
from comfy.k_diffusion.sampling import to_d
from . import filtering
from .utils import fallback
class History:
def __init__(self, size):
self.history = []
self.size = size
def __len__(self):
return len(self.history)
def __getitem__(self, k):
return self.history[k]
def push(self, val):
if len(self.history) >= self.size:
self.history = self.history[-(self.size - 1) :]
self.history.append(val)
def reset(self):
self.history = []
def clone(self):
obj = self.__new__(self.__class__)
obj.__init__(self.size)
obj.history = self.history.copy()
return obj
class ModelResult:
def __init__(
self,
call_idx,
sigma,
x,
denoised,
**kwargs,
):
self.call_idx = call_idx
self.sigma = sigma
self.x = x
self.denoised = denoised
for k in ("denoised_uncond", "denoised_cond", "tangents", "jdenoised"):
setattr(self, k, kwargs.pop(k, None))
if len(kwargs) != 0:
raise ValueError(f"Unexpected keyword arguments: {tuple(kwargs.keys())}")
def to_d(
self,
/,
x=None,
sigma=None,
denoised=None,
denoised_uncond=None,
alt_cfgpp_scale=0,
cfgpp=False,
):
x = fallback(x, self.x)
sigma = fallback(sigma, self.sigma)
denoised = fallback(denoised, self.denoised)
denoised_uncond = fallback(denoised_uncond, self.denoised_uncond)
if alt_cfgpp_scale != 0:
x = x - denoised * alt_cfgpp_scale + denoised_uncond * alt_cfgpp_scale
return to_d(x, sigma, denoised if not cfgpp else denoised_uncond)
@property
def d(self):
return self.to_d()
def clone(self, deep=False):
obj = self.__new__(self.__class__)
for k in (
"denoised",
"call_idx",
"sigma",
"x",
"denoised_uncond",
"denoised_cond",
"tangents",
"jdenoised",
):
val = getattr(self, k)
if deep and isinstance(val, torch.Tensor):
val = val.copy()
setattr(obj, k, val)
return obj
ModelCallCacheConfig = namedtuple(
"ModelCallCacheConfig", ("size", "max_use", "threshold"), defaults=(0, 1000000, 1)
)
class ModelCallCache:
def __init__(
self,
model,
x,
s_in,
extra_args,
*,
cache=None,
filter=None,
):
self.cache = ModelCallCacheConfig(**fallback(cache, {}))
filtargs = fallback(filter, {}).copy()
self.filters = {}
for key in ("input", "denoised", "jdenoised", "cond", "uncond"):
filt = filtargs.pop(key, None)
if filt is None:
continue
self.filters[key] = filtering.make_filter(filt)
self.model = model
self.s_in = s_in
self.extra_args = extra_args
if self.cache.size < 1:
return
self.reset_cache()
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, *args, **kwargs):
if not self.filters:
return result
result = result.clone()
for key in ("denoised", "cond", "uncond", "jdenoised"):
filt = self.filters.get(key)
if filt is None:
continue
attk = f"denoised_{key}" if key in ("cond", "uncond") else key
inpval = getattr(result, attk, None)
if inpval is None:
continue
setattr(result, attk, filt.apply(inpval, *args, **kwargs))
return result
@staticmethod
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 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
def __call__(
self,
x,
sigma,
*,
call_index=0,
ss,
s_in=None,
tangents=None,
return_cached=False,
**kwargs,
):
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
def postcfg(args):
nonlocal denoised_cond, denoised_uncond
denoised_uncond = args["uncond_denoised"]
denoised_cond = args["cond_denoised"]
return args["denoised"]
extra_args = self.extra_args | {
"model_options": comfy.model_patcher.set_model_options_post_cfg_function(
model_options, postcfg, disable_cfg1_optimization=True
)
}
s_in = fallback(s_in, self.s_in)
x = self.maybe_filter("input", x, refs=filter_refs)
def call_model(x, sigma, **kwargs):
return self.model(x, sigma * s_in, **extra_args | kwargs)
if tangents is None:
denoised = call_model(x, sigma, **kwargs)
mr = ModelResult(
call_index,
sigma,
x,
denoised,
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, False) if return_cached else mr
denoised, denoised_prime = torch.func.jvp(call_model, (x, sigma), tangents)
mr = ModelResult(
call_index,
sigma,
x,
denoised,
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, False) if return_cached else mr
+376 -84
View File
@@ -1,14 +1,23 @@
from .sampling import composable_sampler, STEP_SAMPLERS
from .substep_sampling import StepSamplerChain
from .substep_merging import MERGE_SUBSTEPS_CLASSES
import comfy
import yaml
import comfy
class ComposableSampler:
from .sampling import composable_sampler
from .substep_sampling import StepSamplerChain, StepSamplerGroups, ParamGroup
from .step_samplers import STEP_SAMPLERS
from .substep_merging import MERGE_SUBSTEPS_CLASSES
from .restart import Restart
DEFAULT_YAML_PARAMS = """\
# JSON or YAML parameters
s_noise: 1.0
eta: 1.0
"""
class SamplerNode:
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling/samplers"
CATEGORY = "sampling/custom_sampling/OCS"
FUNCTION = "go"
@@ -16,34 +25,17 @@ class ComposableSampler:
def INPUT_TYPES(cls):
return {
"required": {
"s_noise": (
"FLOAT",
{
"default": 1.0,
"min": -100.0,
"max": 100.0,
"step": 0.01,
"round": False,
},
),
"eta": (
"FLOAT",
{
"default": 1.0,
"min": -100.0,
"max": 100.0,
"step": 0.01,
"round": False,
},
),
"merge_method": (tuple(MERGE_SUBSTEPS_CLASSES.keys()),),
"step_sampler_chain": ("STEP_SAMPLER_CHAIN",),
"groups": ("OCS_GROUPS",),
},
"optional": {
"merge_sampler_opt": ("STEP_SAMPLER_CHAIN",),
"params_opt": ("OCS_PARAMS",),
"parameters": (
"STRING",
{"default": "", "multiline": True, "dynamicPrompts": False},
{
"default": DEFAULT_YAML_PARAMS,
"multiline": True,
"dynamicPrompts": False,
},
),
},
}
@@ -51,41 +43,31 @@ class ComposableSampler:
def go(
self,
*,
s_noise,
eta,
merge_method,
step_sampler_chain,
merge_sampler_opt=None,
groups,
params_opt=None,
parameters="",
):
if merge_sampler_opt is not None:
merge_sampler = merge_sampler_opt.items[0]
else:
merge_sampler = ComposableStepSampler().go(step_method="euler")[0].items[0]
options = {
"s_noise": s_noise,
"eta": eta,
"merge_method": merge_method,
"merge_sampler": merge_sampler,
}
options = {}
parameters = parameters.strip()
if parameters:
extra_params = yaml.safe_load(parameters)
if not isinstance(extra_params, dict):
raise ValueError("Parameters must be a JSON or YAML object")
options |= extra_params
options["chain"] = step_sampler_chain.clone()
if extra_params is not None:
if not isinstance(extra_params, dict):
raise ValueError("Parameters must be a JSON or YAML object")
options |= extra_params
if params_opt is not None:
options |= params_opt.items
options["_groups"] = groups.clone()
return (
comfy.samplers.KSAMPLER(
composable_sampler,
{"composable_sampler_options": options},
composable_sampler, {"overly_complicated_options": options}
),
)
class ComposableStepSampler:
RETURN_TYPES = ("STEP_SAMPLER_CHAIN",)
CATEGORY = "sampling/custom_sampling/samplers"
class GroupNode:
RETURN_TYPES = ("OCS_GROUPS",)
CATEGORY = "sampling/custom_sampling/OCS"
FUNCTION = "go"
@@ -93,52 +75,362 @@ class ComposableStepSampler:
def INPUT_TYPES(cls):
return {
"required": {
"s_noise": (
"merge_method": (tuple(MERGE_SUBSTEPS_CLASSES.keys()),),
"time_mode": (("step", "step_pct", "sigma"),),
"time_start": (
"FLOAT",
{
"default": 1.0,
"min": -100.0,
"max": 100.0,
"step": 0.01,
"round": False,
},
{"default": 0, "min": 0.0, "step": 0.1, "round'": False},
),
"eta": (
"time_end": (
"FLOAT",
{
"default": 1.0,
"min": -100.0,
"max": 100.0,
"step": 0.01,
"round": False,
},
{"default": 999, "min": 0.0, "step": 0.1, "round'": False},
),
"substeps": ("INT", {"default": 1, "min": 1, "max": 1000}),
"step_method": (tuple(STEP_SAMPLERS.keys()),),
"substeps": ("OCS_SUBSTEPS",),
},
"optional": {
"step_sampler_opt": ("STEP_SAMPLER_CHAIN",),
"custom_noise_opt": ("SONAR_CUSTOM_NOISE",),
"groups_opt": ("OCS_GROUPS",),
"params_opt": ("OCS_PARAMS",),
"parameters": (
"STRING",
{"default": "", "multiline": True, "dynamicPrompts": False},
{
"default": DEFAULT_YAML_PARAMS,
"multiline": True,
"dynamicPrompts": False,
},
),
},
}
def go(self, *, parameters="", step_sampler_opt=None, **kwargs):
if step_sampler_opt is not None:
chain = step_sampler_opt.clone()
def go(
self,
*,
merge_method,
time_mode,
time_start,
time_end,
substeps,
groups_opt=None,
params_opt=None,
parameters="",
):
group = StepSamplerGroups() if groups_opt is None else groups_opt.clone()
chain = substeps.clone()
chain.merge_method = merge_method
chain.time_mode = time_mode
chain.time_start, chain.time_end = time_start, time_end
options = {}
parameters = parameters.strip()
if parameters:
extra_params = yaml.safe_load(parameters)
if extra_params is not None:
if not isinstance(extra_params, dict):
raise ValueError("Parameters must be a JSON or YAML object")
options |= extra_params
if params_opt is not None:
options |= params_opt.items
chain.options |= options
group.append(chain)
return (group,)
class SubstepsNode:
RETURN_TYPES = ("OCS_SUBSTEPS",)
CATEGORY = "sampling/custom_sampling/OCS"
FUNCTION = "go"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"substeps": ("INT", {"default": 1, "min": 1, "max": 1000}),
"step_method": (tuple(STEP_SAMPLERS.keys()),),
},
"optional": {
"substeps_opt": ("OCS_SUBSTEPS",),
"params_opt": ("OCS_PARAMS",),
"parameters": (
"STRING",
{
"default": DEFAULT_YAML_PARAMS,
"multiline": True,
"dynamicPrompts": False,
},
),
},
}
def go(
self,
*,
parameters="",
substeps_opt=None,
params_opt=None,
**kwargs,
):
if substeps_opt is not None:
chain = substeps_opt.clone()
else:
chain = StepSamplerChain()
parameters = parameters.strip()
if parameters:
extra_params = yaml.safe_load(parameters)
if not isinstance(extra_params, dict):
raise ValueError("Parameters must be a JSON or YAML object")
kwargs |= extra_params
chain.items.append(kwargs)
if extra_params is not None:
if not isinstance(extra_params, dict):
raise ValueError("Parameters must be a JSON or YAML object")
kwargs |= extra_params
if params_opt is not None:
kwargs |= params_opt.items
chain.append(kwargs)
return (chain,)
__all__ = ("ComposableStepSampler", "ComposableSampler")
class Wildcard(str):
__slots__ = ()
def __ne__(self, _unused):
return False
class ParamNode:
RETURN_TYPES = ("OCS_PARAMS",)
CATEGORY = "sampling/custom_sampling/OCS"
FUNCTION = "go"
WC = Wildcard("*")
OCS_PARAM_TYPES = {
"custom_noise": lambda v: hasattr(v, "make_noise_sampler"),
"merge_sampler": lambda v: isinstance(v, StepSamplerChain),
"restart_custom_noise": lambda v: hasattr(v, "make_noise_sampler"),
"SAMPLER": lambda _v: True,
}
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"key": (tuple(cls.OCS_PARAM_TYPES.keys()),),
"value": (cls.WC,),
},
"optional": {
"params_opt": ("OCS_PARAMS",),
"parameters": (
"STRING",
{
"default": "# Additional YAML or JSON parameters\n",
"multiline": True,
"dynamicPrompts": False,
},
),
},
}
def go(self, *, key, value, params_opt=None, parameters=""):
if not self.OCS_PARAM_TYPES[key](value):
raise ValueError(f"CSamplerParam: Bad value type for key {key}")
if parameters:
extra_params = yaml.safe_load(parameters)
if extra_params is not None:
if not isinstance(extra_params, dict):
raise ValueError("Parameters must be a JSON or YAML object")
else:
extra_params = None
params = ParamGroup(items={}) if params_opt is None else params_opt.clone()
params[key] = value
if extra_params is not None:
params[f"{key}.params"] = extra_params
return (params,)
class MultiParamNode:
RETURN_TYPES = ("OCS_PARAMS",)
CATEGORY = "sampling/custom_sampling/OCS"
FUNCTION = "go"
PARAM_COUNT = 5
@classmethod
def INPUT_TYPES(cls):
param_keys = (("", *ParamNode.OCS_PARAM_TYPES.keys()),)
return {
"required": {
f"key_{idx}": param_keys for idx in range(1, cls.PARAM_COUNT + 1)
},
"optional": {
"params_opt": ("OCS_PARAMS",),
"parameters": (
"STRING",
{
"default": """\
# Additional YAML or JSON parameters
# Should be an object with key corresponding to the index of the input
""",
"multiline": True,
"dynamicPrompts": False,
},
),
}
| {
f"value_opt_{idx}": (ParamNode.WC,)
for idx in range(1, cls.PARAM_COUNT + 1)
},
}
def go(self, *, params_opt=None, parameters="", **kwargs):
params = ParamGroup(items={}) if params_opt is None else params_opt.clone()
if parameters:
extra_params = yaml.safe_load(parameters)
if extra_params is not None:
if not isinstance(extra_params, dict):
raise ValueError("Parameters must be a JSON or YAML object")
else:
extra_params = {}
else:
extra_params = {}
for idx in range(1, self.PARAM_COUNT + 1):
key, value = kwargs.get(f"key_{idx}"), kwargs.get(f"value_opt_{idx}")
if not key or value is None:
continue
if not ParamNode.OCS_PARAM_TYPES[key](value):
raise ValueError(f"CSamplerParamGroup: Bad value type for key {key}")
params[key] = value
extra = extra_params.get(str(idx))
if extra is not None:
params[f"{key}.params"] = extra
return (params,)
class SimpleRestartSchedule:
RETURN_TYPES = ("SIGMAS",)
CATEGORY = "sampling/custom_sampling/OCS"
FUNCTION = "go"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"sigmas": ("SIGMAS",),
"start_step": ("INT", {"min": 0, "default": 0}),
},
"optional": {
"schedule": (
"STRING",
{
"default": """\
# YAML or JSON restart schedule
# Every 5 steps, jump back 3 steps
- [5, -3]
# Jump to schedule item 0
- 0
""",
"multiline": True,
"dynamicPrompts": False,
},
),
},
}
def go(self, *, sigmas, start_step=0, schedule="[]"):
if schedule:
parsed_schedule = yaml.safe_load(schedule)
if parsed_schedule is not None:
if not isinstance(parsed_schedule, (list, tuple)):
raise ValueError("Schedule must be a JSON or YAML list")
else:
parsed_schedule = []
else:
parsed_schedule = []
return (Restart.simple_schedule(sigmas, start_step, parsed_schedule),)
class ModelSetMaxSigmaNode:
RETURN_TYPES = ("MODEL",)
CATEGORY = "hacks"
FUNCTION = "go"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"mode": (("recalculate", "simple_multiply"),),
"sigma_max": (
"FLOAT",
{
"default": -1.0,
"min": -10000.0,
"max": 10000.0,
"step": 0.01,
"round'": False,
},
),
"fake_sigma_min": (
"FLOAT",
{
"default": 0.0,
"min": 0.0,
"max": 1000.0,
"step": 0.01,
"round'": False,
},
),
}
}
def go(self, model, mode="recalculate", sigma_max=-1.0, fake_sigma_min=0.0):
if sigma_max == 0:
raise ValueError("ModelSetMaxSigma: Invalid sigma_max value")
if mode not in ("recalculate", "simple_multiply"):
raise ValueError("ModelSetMaxSigma: Invalid mode value")
orig_ms = model.get_model_object("model_sampling")
model = model.clone()
orig_max_sigma, orig_min_sigma = (
orig_ms.sigma_max.item(),
orig_ms.sigma_min.item(),
)
max_multiplier = abs(sigma_max) if sigma_max < 0 else sigma_max / orig_max_sigma
if max_multiplier == 1:
return (model,)
mcfg = model.get_model_object("model_config")
orig_sigmas = orig_ms.sigmas
fake_sigma_min = orig_sigmas.new_full((1,), fake_sigma_min)
class NewModelSampling(orig_ms.__class__):
if fake_sigma_min != 0:
@property
def sigma_min(self):
return fake_sigma_min
ms = NewModelSampling(mcfg)
if mode == "simple_multiply":
ms.set_sigmas(orig_sigmas * max_multiplier)
else:
ss = getattr(mcfg, "sampling_setting", None) or {}
if ss.get("beta_schedule", "linear") != "linear":
raise NotImplementedError(
"ModelSetMaxSigma: Can only handle linear beta schedules in reschedule mode"
)
ms.set_sigmas((orig_sigmas**2 * max_multiplier**2) ** 0.5)
new_max_sigma, new_min_sigma = ms.sigma_max.item(), ms.sigma_min.item()
if new_min_sigma >= new_max_sigma:
raise ValueError(
"ModelSetMaxSigma: Invalid fake_min_sigma value, result max <= min"
)
model.add_object_patch("model_sampling", ms)
print(
f"ModelSetMaxSigma: Set model sigmas({mode}): old_max={orig_max_sigma:.04}, old_min={orig_min_sigma:.03}, new_max={new_max_sigma:.04}, new_min={new_min_sigma:.03}"
)
return (model,)
__all__ = (
"SamplerNode",
"GroupNode",
"SubstepsNode",
"ParamNode",
"MultiParamNode",
"ModelSetMaxSigmaNode",
)
+231
View File
@@ -0,0 +1,231 @@
import gc
import random
import scipy
import torch
from .filtering import Filter, make_filter
from .utils import scale_noise, fallback
class ImmiscibleNoise(Filter):
name = "immiscible"
uses_ref = True
default_options = Filter.default_options | {
"size": 0,
"batching": "channel",
"maximize": False,
}
def __call__(self, noise_sampler, x_ref, *, refs=None):
if not self.check_applies(refs):
return noise_sampler()
return self.apply(
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,
)
def filter(self, latent, ref_latent, *, refs, output_shape):
if self.size == 0:
return latent
return self.unbatch(
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.view(sz[0] * sz[1], *sz[2:])
if self.batching == "row":
return latent.view(sz[0] * sz[1] * sz[2], sz[3])
if self.batching == "column":
return latent.permute(0, 1, 3, 2).reshape(sz[0] * sz[1] * sz[3], sz[2])
raise ValueError("Bad Immmiscible noise batching type")
def unbatch(self, latent, sz):
if self.batching == "column":
return latent.view(*sz[:2], sz[3], sz[2]).permute(0, 1, 3, 2)
return latent.view(*sz)
# 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
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:
def __init__(
self,
x,
seed,
min_sigma,
max_sigma,
*,
normalize_noise=True,
cpu_noise=True,
batch_size=32,
caching=True,
cache_reset_interval=1,
set_seed=False,
scale=1.0,
normalize_dims=(-3, -2, -1),
immiscible=None,
filter=None,
**_unused,
):
self.x = x
self.mega_x = None
self.seed = seed
self.seed_offset = 0
self.min_sigma = min_sigma
self.max_sigma = max_sigma
self.cache = {}
self.batch_size = max(1, batch_size)
self.normalize_noise = normalize_noise
self.cpu_noise = cpu_noise
self.caching = caching
self.cache_reset_interval = max(1, cache_reset_interval)
self.scale = float(scale)
self.normalize_dims = tuple(int(v) for v in normalize_dims)
self.immiscible = ImmiscibleNoise(**fallback(immiscible, {}))
if filter is None:
self.filter = None
else:
self.filter = make_filter(filter)
self.update_x(x)
if set_seed:
random.seed(seed)
torch.manual_seed(seed)
def reset_cache(self):
self.cache = {}
gc.collect()
def scale_noise(self, noise, factor=1.0, normalized=None, normalize_dims=None):
normalized = self.normalize_noise if normalized is None else normalized
normalize_dims = (
self.normalize_dims if normalize_dims is None else normalize_dims
)
return scale_noise(
noise, factor, normalized=normalized, normalize_dims=normalize_dims
)
def update_x(self, x):
if self.x.shape == x.shape and self.mega_x is not None:
self.x = x
return
self.x = x
self.mega_x = None
self.reset_cache()
if self.batch_size == 1:
self.mega_x = x
return
self.mega_x = x.repeat(x.shape[0] * self.batch_size, *((1,) * (x.dim() - 1)))
def set_cache(self, key, noise_sampler):
if not self.caching:
return
self.cache[key] = noise_sampler
def make_caching_noise_sampler(
self,
nsobj,
size,
sigma,
sigma_next,
immiscible=None,
):
size = min(size, self.batch_size)
cache_key = (nsobj, size)
if self.caching:
noise_sampler = self.cache.get(cache_key)
if noise_sampler:
return noise_sampler
curr_seed = self.seed + self.seed_offset
self.seed_offset += 1
curr_x = self.mega_x[: self.x.shape[0] * size, ...]
if nsobj is None:
def ns(_s, _sn, *_unused, **_unusedkwargs):
return torch.randn_like(curr_x)
else:
ns = nsobj.make_noise_sampler(
curr_x,
self.min_sigma,
self.max_sigma,
seed=curr_seed,
normalized=False,
cpu=self.cpu_noise,
)
orig_h, orig_w = self.x.shape[-2:]
remain = 0
noise = None
if immiscible is None:
immiscible = self.immiscible
def noise_sampler_(
*_unused,
out_hw=(orig_h, orig_w),
**_unusedkwargs,
):
nonlocal remain, noise
if out_hw != (orig_h, orig_w):
raise NotImplementedError(
f"Noise size mismatch: {out_hw} vs {(orig_h, orig_w)}"
)
if remain < 1:
noise = self.scale_noise(ns(sigma, sigma_next)).view(
size,
*self.x.shape,
)
remain = size
result = noise[-remain]
remain -= 1
return result
def noise_sampler(*args, x_ref=None, refs=None, **kwargs):
if immiscible is False:
noise = noise_sampler_(*args, **kwargs)
else:
noise = immiscible(
lambda args=args, kwargs=kwargs: noise_sampler_(*args, **kwargs),
fallback(x_ref, self.x),
refs=refs,
)
return (
self.filter.apply(noise, refs=refs)
if self.filter is not None
else noise
)
self.set_cache(cache_key, noise_sampler)
return noise_sampler
+91
View File
@@ -0,0 +1,91 @@
import torch
class Restart:
def __init__(self, *, s_noise=1.0, custom_noise=None, immiscible=False):
from .noise import ImmiscibleNoise
self.s_noise = s_noise
if immiscible is not False:
immiscible = ImmiscibleNoise(**immiscible)
self.immiscible = immiscible
self.custom_noise = custom_noise
def get_noise_sampler(self, nsc):
return nsc.make_caching_noise_sampler(
self.custom_noise,
1,
nsc.max_sigma,
nsc.min_sigma,
immiscible=self.immiscible,
)
@staticmethod
def get_segment(sigmas: torch.Tensor) -> torch.Tensor:
last_sigma = sigmas[0]
for idx in range(1, len(sigmas)):
sigma = sigmas[idx]
if sigma > last_sigma:
return sigmas[:idx]
last_sigma = sigma
return sigmas
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]:
noise_scale = self.get_noise_scale(prev_seg[-1], seg[0])
else:
noise_scale = 0.0
prev_seg = seg
yield (noise_scale, seg)
def get_noise_scale(self, s_min, s_max):
result = (s_max**2 - s_min**2) ** 0.5
if isinstance(result, torch.Tensor):
result = result.item()
return result * self.s_noise
@classmethod
def simple_schedule(cls, sigmas, start_step, schedule=(), max_iter=1000):
if sigmas.ndim != 1:
raise ValueError("Bad number of dimensions for sigmas")
siglen = len(sigmas) - 1
if siglen <= start_step or not len(schedule):
return sigmas
siglist = sigmas.cpu().tolist()
out = siglist[:start_step]
sched_len = len(schedule)
sched_idx = 0
sig_idx = start_step
iter_count = 0
while 0 <= sched_idx < sched_len:
# print(f"LOOP: sched_idx={sched_idx}, sig_idx={sig_idx}: {out}")
iter_count += 1
if iter_count > max_iter:
raise RuntimeError("Hit max iteration count. Loop in schedule?")
item = schedule[sched_idx]
if not isinstance(item, (list, tuple)):
if item < 0:
item = sched_len + item
if item < 0 or item >= sched_len:
raise ValueError("Schedule jump index out of range")
sched_idx = item
continue
if sig_idx >= siglen or sig_idx < 0:
break
interval, jump = item
chunk = siglist[sig_idx : sig_idx + interval + 1]
# print(f"{out} + {chunk}")
out += chunk
sig_idx += interval + jump
if jump >= 0:
sig_idx += 1
sched_idx += 1
if sig_idx < siglen and sig_idx >= 0:
out += siglist[sig_idx:]
if out[-1] > siglist[-1]:
out.append(siglist[-1])
return torch.tensor(out).to(sigmas)
+78 -47
View File
@@ -2,9 +2,22 @@ import torch
from tqdm.auto import trange
from .substep_samplers import STEP_SAMPLERS
from .substep_sampling import SamplerState, History, ModelCallCache
from .filtering import FILTER_HANDLERS
from .model import ModelCallCache
from .noise import NoiseSamplerCache
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_with_refs(ss.refs)
if merge_sampler.check_match(handlers, ss=ss):
return merge_sampler
return None
def composable_sampler(
@@ -14,14 +27,14 @@ def composable_sampler(
*,
s_noise=1.0,
eta=1.0,
composable_sampler_options,
overly_complicated_options,
extra_args=None,
callback=None,
disable=None,
noise_sampler=None,
**kwargs,
):
copts = composable_sampler_options.copy()
copts = overly_complicated_options.copy()
if extra_args is None:
extra_args = {}
if noise_sampler is None:
@@ -29,64 +42,82 @@ def composable_sampler(
def noise_sampler(_s, _sn):
return torch.randn_like(x)
samplers = []
substeps = 0
for sitem in copts["chain"].items:
custom_noise = sitem.get("custom_noise_opt")
if custom_noise is None:
curr_ns = noise_sampler
else:
curr_ns = custom_noise.make_noise_sampler(
x, sigmas[-1], sigmas[0], normalized=True
)
ssampler = STEP_SAMPLERS[sitem["step_method"]](noise_sampler=curr_ns, **sitem)
samplers.append(ssampler)
# samplers += (ssampler,) * sitem["substeps"]
substeps += ssampler.substeps
msitem = copts["merge_sampler"]
if copts["merge_method"] in ("sample", "sample_uncached"):
custom_noise = msitem.get("custom_noise_opt")
if custom_noise is None:
curr_ns = noise_sampler
else:
curr_ns = custom_noise.make_noise_sampler(
x, sigmas[-1], sigmas[0], normalized=True
)
merge_sampler = STEP_SAMPLERS[msitem["step_method"]](
noise_sampler=curr_ns, **msitem
)
pass
else:
merge_sampler = None
restart_params = copts.get("restart", {})
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(
ModelCallCache(
model,
x,
x.new_ones((x.shape[0],)),
extra_args,
size=copts.get("model_call_cache", 0),
max_use=copts.get("model_call_cache_max_use", 1000000),
threshold=copts.get("model_call_cache_threshold", 0),
**copts.get("model", {}),
),
sigmas,
0,
History(x, 3),
History(x, 2),
extra_args,
noise_sampler=noise_sampler,
callback=callback,
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,
)
merge_sampler = MERGE_SUBSTEPS_CLASSES[copts["merge_method"]](
ss,
samplers,
**(copts | {"merge_sampler": merge_sampler}),
groups = copts["_groups"]
merge_samplers = tuple(
MERGE_SUBSTEPS_CLASSES[g.merge_method](ss, g) for g in groups.items
)
for idx in trange(len(sigmas) - 1, disable=disable):
print(f"STEP {idx+1}")
ss.update(idx)
ss.model.reset_cache()
x = merge_sampler.step(x)
nsc = NoiseSamplerCache(
x,
extra_args.get("seed", 42),
sigmas[-1],
sigmas[0],
**copts.get("noise", {}),
)
ss.noise = nsc
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 noise_scale, chunk_sigmas in sigma_chunks:
ss.sigmas = chunk_sigmas
ss.update(0, step=step, substep=0)
if step != 0:
nsc.reset_cache()
ss.hist.reset()
for ms in merge_samplers:
ms.reset()
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 restart_ns
for idx in range(len(chunk_sigmas) - 1):
if idx > 0:
ss.update(idx, step=step, substep=0)
# 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:
raise RuntimeError(f"No matching sampler group for step {step + 1}")
pbar.set_description(
f"{merge_sampler.name}: {ss.sigma.item():.03} -> {ss.sigma_next.item():.03}"
)
x = merge_sampler(x)
if (idx + 1) % nsc.cache_reset_interval == 0:
nsc.reset_cache()
step += 1
pbar.update(1)
return x
+2157
View File
File diff suppressed because it is too large Load Diff
+460 -196
View File
@@ -1,202 +1,363 @@
import torch
import operator
from .utils import scale_noise, find_first_unsorted
from .substep_sampling import History
import torch
import tqdm
from . import expression as expr
from . import utils
from .filtering import make_filter, FilterRefs
from .restart import Restart
from .step_samplers import STEP_SAMPLERS
from .utils import check_time, fallback
class MergeSubstepsSampler:
def __init__(self, ss, samplers, **_kwargs):
name = "unknown"
def __init__(self, ss, group):
samplers = tuple(
STEP_SAMPLERS[sitem["step_method"]](**sitem) for sitem in group.items
)
options = group.options.copy()
self.time_mode = group.time_mode
self.time_start = group.time_start
self.time_end = group.time_end
self.ss = ss
self.samplers = samplers
self.substeps = sum(sampler.substeps for sampler in samplers)
when_expr = options.pop("when", None)
self.when = expr.Expression(when_expr) if when_expr else None
pre_filter = options.pop("pre_filter", None)
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.options = options
def check_match(self, handlers: None | object, *, ss: None | object = None):
ss = fallback(ss, self.ss)
if not check_time(
self.time_mode,
self.time_start,
self.time_end,
ss.sigma,
ss.step,
ss.total_steps,
):
return False
if self.when is None:
return True
if handlers is None:
raise ValueError("Group has when expression but handlers not passed")
return operator.truth(self.when.eval(handlers))
def step_input(self, x, *, ss=None):
if self.pre_filter is None:
return x
ss = fallback(ss, self.ss)
return self.pre_filter.apply(x, refs=fallback(ss, self.ss).refs)
def step_output(self, x, *, orig_x=None, ss=None):
if self.post_filter is None:
return x
ss = fallback(ss, self.ss)
refs = ss.refs if orig_x is None else ss.refs | FilterRefs({"orig_x": orig_x})
return self.post_filter.apply(x, refs=refs)
def __call__(self, x):
orig_x = x
x = self.step_input(x)
x = self.step(x)
return self.step_output(x, orig_x=orig_x)
def step(self, x):
raise NotImplementedError
def merge_steps(self, _x, result):
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, ss=None):
for sr in self.substep(x, sampler, ss=ss):
if not sr.final:
sr.noise_x(ss=fallback(ss, self.ss))
return sr
def merge_steps(self, x, result=None, *, noise=None, ss=None, denoised=True):
ss = ss if ss is not None else self.ss
result = fallback(result, x)
if noise is not None:
result = result + noise
return result
def step_max_noise_samples(self):
return sum(
(1 + sampler.self_noise) * sampler.substeps for sampler in self.samplers
)
def reset(self):
pass
class SimpleSubstepsSampler(MergeSubstepsSampler):
name = "simple"
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
if not len(self.samplers):
raise ValueError("Missing sampler")
def step_max_noise_samples(self):
return 1 + self.samplers[0].self_noise
def step(self, x):
ss, ssampler = self.ss, self.samplers[0]
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)
ss.callback()
sr = self.simple_substep(x, ssampler)
return self.merge_steps(sr.x, noise=sr.get_noise(ss=ss))
class NormalMergeSubstepsSampler(MergeSubstepsSampler):
def __init__(self, ss, samplers, **kwargs):
super().__init__(ss, samplers, **kwargs)
self.ss = ss
name = "normal"
def step(self, x):
ss = self.ss
substeps = self.substeps
renoise_weight = 1.0 / substeps
z_avg = torch.zeros_like(x)
noise = torch.zeros_like(x)
noise = z_avg.clone()
noise_total = 0.0
for idx, ssampler in enumerate(
sampler for sampler in self.samplers for _ in range(sampler.substeps)
):
print(f" SUBSTEP {idx+1}: {ssampler.name}")
ss.denoised = ss.model(x, ss.sigma)
z_k, noise_strength = ssampler.step(x, ss)
z_avg += renoise_weight * z_k
noise_strength *= ssampler.s_noise
if ss.sigma_next == 0 or noise_strength == 0:
continue
noise_curr = ssampler.noise_sampler(ss.sigma, ss.sigma_next)
x = z_k
if idx != substeps - 1:
x += noise_curr * noise_strength
noise_total += noise_strength.item() * renoise_weight
noise += noise_curr * noise_strength
ss.dhist.push(ss.denoised)
ss.denoised = None
x = self.merge_steps(x, z_avg)
if ss.sigma_next != 0 and noise_total != 0:
x += scale_noise(noise, noise_total * ss.s_noise)
ss.xhist.push(x)
ss.callback(x)
return x
class AverageMergeSubstepsSampler(NormalMergeSubstepsSampler):
def __init__(self, ss, samplers, *, avgmerge_stretch=0.4, **kwargs):
super().__init__(ss, samplers, **kwargs)
self.ss = ss
self.stretch = avgmerge_stretch
def step(self, x):
ss = orig_ss = self.ss
substeps = self.substeps
renoise_weight = 1.0 / substeps
z_avg = torch.zeros_like(x)
noise = torch.zeros_like(x)
stretch = (ss.sigma - ss.sigma_next) * self.stretch
sig_adj = ss.sigma + stretch
ss = self.ss.clone_edit(sigma=sig_adj)
orig_x = x
x = x + ss.noise_sampler(orig_ss.sigma, ss.sigma_next) * stretch * ss.s_noise
ss.denoised = ss.model(x, sig_adj)
noise_total = 0.0
step = 0
for idx, ssampler in enumerate(self.samplers):
print(
f" SUBSTEP {step+1} .. {step+ssampler.substeps}: {ssampler.name}, stretch={stretch}"
substep = 0
pbar = tqdm.tqdm(total=self.substeps, initial=1, disable=ss.disable_status)
ss.hist.push(ss.model(x, ss.sigma, ss=ss))
ss.refs = FilterRefs.from_ss(ss, have_current=True)
ss.callback()
for ssampler in self.samplers:
custom_noise = ssampler.options.get(
"custom_noise", self.options.get("custom_noise")
)
for sidx in range(ssampler.substeps):
curr_x = orig_x + scale_noise(
ssampler.noise_sampler(sig_adj, ss.sigma_next), stretch
)
z_k, noise_strength = ssampler.step(curr_x, ss)
z_avg += renoise_weight * z_k
if ss.sigma_next == 0:
continue
noise_strength *= ssampler.s_noise
if noise_strength == 0:
continue
noise_curr = ssampler.noise_sampler(ss.sigma, ss.sigma_next)
noise_total += noise_strength.item() * renoise_weight
noise += noise_curr * noise_strength
step += ssampler.substeps
ss.dhist.push(ss.denoised)
ss.denoised = None
x = self.merge_steps(x, z_avg)
if ss.sigma_next != 0 and noise_total != 0:
x += scale_noise(noise, noise_total * ss.s_noise)
ss.xhist.push(x)
ss.callback(x)
return x
class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler):
cache_model = True
def __init__(self, ss, samplers, *, merge_sampler, **kwargs):
super().__init__(ss, samplers, **kwargs)
self.merge_sampler = merge_sampler
self.merge_ss = None
def step(self, x):
ss = self.ss
substeps = self.substeps
renoise_weight = 1.0 / substeps
z_avg = torch.zeros_like(x)
curr_x = x
ss.denoised = None
stretch = (ss.sigma - ss.sigma_next) * self.stretch
sig_adj = ss.sigma + stretch
ss = self.ss.clone_edit(sigma=sig_adj)
step = 0
for idx, ssampler in enumerate(self.samplers):
print(
f" SUBSTEP {step+1} .. {step+ssampler.substeps}: {ssampler.name}, stretch={stretch}"
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),
)
if idx == 0 or not self.cache_model:
ss.denoised = ss.model(
curr_x,
# + ss.noise_sampler(sig_adj.sigma, ss.sigma_next) * stretch * ss.s_noise,
sig_adj,
)
for sidx in range(ssampler.substeps):
curr_x = (
x
+ ssampler.noise_sampler(sig_adj, ss.sigma_next)
* ssampler.s_noise
* stretch
)
z_k, noise_strength = ssampler.step(curr_x, ss)
z_avg += renoise_weight * z_k
curr_x = z_k
if noise_strength == 0 or ss.sigma_next == 0:
continue
curr_x += (
ssampler.noise_sampler(ss.sigma, ss.sigma_next)
* ssampler.s_noise
* noise_strength
)
step += ssampler.substeps
ss.dhist.push(ss.denoised)
ss.denoised = None
x = self.merge_steps(curr_x, z_avg)
ss.xhist.push(x)
ss.callback(x)
return x
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)
def merge_steps(self, x, result):
self.ss.model.reset_cache()
msampler = self.merge_sampler
if self.merge_ss is None:
merge_ss = self.merge_ss = self.ss.clone_edit(
denoised=result,
dhist=History(x, 3),
xhist=History(x, 2),
s_noise=msampler.s_noise,
eta=msampler.eta,
# model_call_cache=None,
)
else:
merge_ss = self.merge_ss
merge_ss.denoised = result
merge_ss.update(self.ss.idx)
final = merge_ss.sigma_next == 0
merged, noise_strength = msampler.step(x, merge_ss)
if not final:
ss = self.ss
merged = (
merged
+ msampler.noise_sampler(ss.sigma, ss.sigma_next)
* msampler.s_noise
* ss.sigma_up
)
merge_ss.dhist.push(result)
merge_ss.xhist.push(merged)
merge_ss.denoised = None
return merged
noise = ss.noise.scale_noise(
noise,
noise_total * self.options.get("s_noise", 1.0),
normalized=True,
)
return self.merge_steps(
x, z_avg, noise=None if noise_total == 0 else noise, denoised=ss.denoised
)
class SampleUncachedMergeSubstepsSampler(SampleMergeSubstepsSampler):
cache_model = False
# class AverageMergeSubstepsSampler(NormalMergeSubstepsSampler):
# name = "average"
# def __init__(self, ss, sitems, *, avgmerge_stretch=0.4, **kwargs):
# super().__init__(ss, sitems, **kwargs)
# self.stretch = avgmerge_stretch
# def step_max_noise_samples(self):
# return sum(
# 1 + (2 + sampler.self_noise) * sampler.substeps for sampler in self.samplers
# )
# def step(self, x):
# ss = orig_ss = self.ss
# substeps = self.substeps
# renoise_weight = 1.0 / substeps
# z_avg = torch.zeros_like(x)
# noise = torch.zeros_like(x)
# stretch = (ss.sigma - ss.sigma_next) * self.stretch
# sig_adj = ss.sigma + stretch
# ss = self.ss.clone_edit(sigma=sig_adj)
# orig_x = x
# stretch_strength = stretch * ss.s_noise
# if stretch_strength != 0:
# noise_sampler = ss.noise.make_caching_noise_sampler(
# self.options.get("custom_noise"), 1, orig_ss.sigma, ss.sigma_next
# )
# x = x + (
# noise_sampler(orig_ss.sigma, ss.sigma_next).mul_(stretch * ss.s_noise)
# )
# self.ss.denoised = ss.denoised = ss.model(x, sig_adj)
# noise_total = 0.0
# substep = 0
# for idx, ssampler in enumerate(self.samplers):
# print(
# f" SUBSTEP {substep + 1} .. {substep + ssampler.substeps}: {ssampler.name}, stretch={stretch}"
# )
# custom_noise = ssampler.options.get(
# "custom_noise", self.options.get("custom_noise")
# )
# noise_sampler = ss.noise.make_caching_noise_sampler(
# custom_noise,
# ssampler.substeps
# + (0 if ss.sigma_next == 0 else ssampler.max_noise_samples()),
# ss.sigma,
# ss.sigma_next,
# )
# ssampler.noise_sampler = noise_sampler
# for sidx in range(ssampler.substeps):
# curr_x = orig_x + noise_sampler(sig_adj, ss.sigma_next).mul_(stretch)
# sr = self.simple_substep(curr_x, ssampler, ss=ss)
# z_avg += renoise_weight * sr.x
# noise_strength = sr.noise_scale
# if ss.sigma_next == 0 or noise_strength == 0:
# continue
# if noise_strength != 0 and ss.sigma_next != 0:
# noise_curr = sr.get_noise()
# noise_total += noise_strength.item() * renoise_weight
# noise += noise_curr
# substep += 1
# substep += ssampler.substeps
# return self.merge_steps(
# x,
# z_avg,
# noise=None
# if not noise_total
# else ss.noise.scale_noise(noise, noise_total * ss.s_noise, normalized=True),
# ss=ss,
# )
# class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler):
# name = "sample"
# cache_model = True
# def __init__(self, ss, sitems, *, merge_sampler=None, **kwargs):
# super().__init__(ss, sitems, **kwargs)
# if merge_sampler is None:
# merge_sampler = STEP_SAMPLERS["euler"](step_method="euler")
# else:
# msitem = merge_sampler.items[0]
# merge_sampler = STEP_SAMPLERS[msitem["step_method"]](**msitem)
# self.merge_sampler = merge_sampler
# self.merge_ss = None
# def step(self, x):
# ss = self.ss
# substeps = self.substeps
# renoise_weight = 1.0 / substeps
# z_avg = torch.zeros_like(x)
# curr_x = x
# ss.denoised = None
# stretch = (ss.sigma - ss.sigma_next) * self.stretch
# sig_adj = ss.sigma + stretch
# ss = self.ss.clone_edit(sigma=sig_adj)
# step = 0
# for idx, ssampler in enumerate(self.samplers):
# print(
# f" SUBSTEP {step + 1} .. {step + ssampler.substeps}: {ssampler.name}, stretch={stretch}"
# )
# 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() + ssampler.substeps,
# ss.sigma,
# ss.sigma_next,
# )
# ssampler.noise_sampler = noise_sampler
# for sidx in range(ssampler.substeps):
# if idx + sidx == 0 or not self.cache_model:
# self.ss.denoised = ss.denoised = ss.model(
# curr_x,
# ss.sigma,
# # + ss.noise_sampler(sig_adj.sigma, ss.sigma_next) * stretch * ss.s_noise,
# # sig_adj,
# )
# curr_x = x + noise_sampler(sig_adj, ss.sigma_next).mul_(
# ssampler.s_noise * stretch
# )
# sr = self.simple_substep(curr_x, ssampler, ss=ss)
# z_avg += renoise_weight * sr.x
# curr_x = sr.noise_x(sr.x)
# step += ssampler.substeps
# return self.merge_steps(curr_x, z_avg)
# def merge_steps(self, x, result):
# ss = self.ss
# ss.dhist.push(ss.denoised)
# ss.denoised = None
# ss.model.reset_cache()
# msampler = self.merge_sampler
# if self.merge_ss is None:
# merge_ss = self.merge_ss = self.ss.clone_edit(
# denoised=result,
# dhist=History(x, 3),
# xhist=History(x, 2),
# s_noise=msampler.s_noise,
# eta=msampler.eta,
# )
# else:
# merge_ss = self.merge_ss
# merge_ss.denoised = result
# merge_ss.update(self.ss.idx, step=self.ss.step)
# 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),
# merge_ss.sigma,
# merge_ss.sigma_next,
# )
# msampler.noise_sampler = noise_sampler
# sr = self.simple_substep(x, msampler, ss=merge_ss)
# self.ss.callback(sr.x)
# sr.noise_x()
# merge_ss.dhist.push(result)
# merge_ss.xhist.push(sr.x)
# merge_ss.denoised = None
# ss.xhist.push(sr.x)
# return sr.x
# def reset(self):
# if self.merge_ss is None:
# return
# self.merge_ss.reset()
# self.merge_ss.sigmas = self.ss.sigmas
# self.merge_ss.update(self.ss.idx, step=self.ss.step)
# class SampleUncachedMergeSubstepsSampler(SampleMergeSubstepsSampler):
# name = "sample_uncached"
# cache_model = False
class DivideMergeSubstepsSampler(MergeSubstepsSampler):
def __init__(self, ss, samplers, *, schedule_multiplier=4, **kwargs):
super().__init__(ss, samplers, **kwargs)
name = "divide"
def __init__(self, ss, group, *, schedule_multiplier=4, **kwargs):
super().__init__(ss, group, **kwargs)
self.schedule_multiplier = schedule_multiplier
def make_schedule(self, ss):
@@ -204,11 +365,9 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler):
sigmas_slice = ss.sigmas[
ss.idx : min(max_steps + 1, ss.idx + self.schedule_multiplier)
]
# print("SLICE", sigmas_slice)
unsorted_idx = find_first_unsorted(sigmas_slice)
unsorted_idx = utils.find_first_unsorted(sigmas_slice)
if unsorted_idx is not None:
sigmas_slice = sigmas_slice[:unsorted_idx]
# print("SLICE ADJ", sigmas_slice)
chunks = tuple(
torch.linspace(
sigmas_slice[idx],
@@ -219,42 +378,147 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler):
)[0 if not idx else 1 :]
for idx in range(len(sigmas_slice) - 1)
)
# print("CHUNKS", chunks)
return torch.cat(chunks)
def step(self, x):
ss = self.ss
# print("SUBSIGMAS", subsigmas)
subss = self.ss.clone_edit(idx=0, sigmas=self.make_schedule(ss))
subss.main_idx = ss.idx
subss.main_sigmas = ss.sigmas
for idx, ssampler in enumerate(
sampler for sampler in self.samplers for _ in range(sampler.substeps)
):
print(f" SUBSTEP {idx+1}: {ssampler.name}")
subss.update(idx)
subss.denoised = subss.model(x, subss.sigma)
x, noise_strength = ssampler.step(x, subss)
if noise_strength == 0 or subss.sigma_next == 0:
continue
x = (
x
+ ssampler.noise_sampler(subss.sigma, subss.sigma_next)
* ssampler.s_noise
* noise_strength
substep = 0
pbar = tqdm.tqdm(total=self.substeps, initial=0, disable=ss.disable_status)
for ssampler in self.samplers:
custom_noise = ssampler.options.get(
"custom_noise", self.options.get("custom_noise")
)
subss.xhist.push(x)
subss.dhist.push(subss.denoised)
subss.denoised = None
ss.callback(x)
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
class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
name = "overshoot"
def __init__(
self,
ss,
group,
*,
overshoot_expand_steps=1,
restart_custom_noise=None,
restart=None,
**kwargs,
):
super().__init__(ss, group, **kwargs)
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),
)
def make_schedule(self, ss):
expand = self.overshoot_expand_steps
if expand > self.substeps:
raise ValueError(
"overshoot_expand_steps > substeps: can't make it to the end of step 1"
)
if expand < 2:
return ss.sigmas, ss.idx
sigmas_cpu = ss.sigmas.cpu()
sigmas = torch.cat(
tuple(
torch.linspace(f, t, expand + 1)[:-1]
for f, t in torch.stack((sigmas_cpu[:-1], sigmas_cpu[1:]), dim=1)
)
+ (sigmas_cpu[-1].unsqueeze(0),)
)
return sigmas.to(ss.sigmas), ss.idx * expand
def step(self, x):
ss = self.ss
sigmas, sigidx = self.make_schedule(ss)
subss = ss.clone_edit(idx=sigidx, sigmas=sigmas)
subss.hist = subss.hist.clone()
substep = 0
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:
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:
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
MERGE_SUBSTEPS_CLASSES = {
"default (simple)": SimpleSubstepsSampler,
"normal": NormalMergeSubstepsSampler,
"divide": DivideMergeSubstepsSampler,
"average": AverageMergeSubstepsSampler,
"sample": SampleMergeSubstepsSampler,
"sample_uncached": SampleUncachedMergeSubstepsSampler,
"overshoot": OvershootMergeSubstepsSampler,
# "average": AverageMergeSubstepsSampler,
# "sample": SampleMergeSubstepsSampler,
# "sample_uncached": SampleUncachedMergeSubstepsSampler,
"simple": SimpleSubstepsSampler,
}
-624
View File
@@ -1,624 +0,0 @@
import math
import torch
from comfy.k_diffusion.sampling import (
get_ancestral_step,
to_d,
)
from .res_support import _de_second_order
from .utils import find_first_unsorted
class SingleStepSampler:
name = None
def __init__(
self,
*,
noise_sampler=None,
substeps=1,
s_noise=1.0,
eta=1.0,
dyn_eta_start=None,
dyn_eta_end=None,
weight=1.0,
**kwargs,
):
self.s_noise = s_noise
self.eta = eta
self.dyn_eta_start = dyn_eta_start
self.dyn_eta_end = dyn_eta_end
self.noise_sampler = noise_sampler
self.weight = weight
self.substeps = substeps
self.kwargs = kwargs
def step(self, x, ss):
raise NotImplementedError
# Euler - based on original ComfyUI implementation
def euler_step(self, x, ss):
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
d = to_d(x, ss.sigma, ss.denoised)
dt = sigma_down - ss.sigma
return x + d * dt, sigma_up
def __str__(self):
return f"<SS({self.name}): s_noise={self.s_noise}, eta={self.eta}>"
def get_dyn_value(self, ss, start, end):
if None in (start, end):
return 1.0
if start == end:
return start
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, ss):
return self.eta * self.get_dyn_value(ss, self.dyn_eta_start, self.dyn_eta_end)
class ReversibleSingleStepSampler(SingleStepSampler):
def __init__(self, *, reta=1.0, dyn_reta_start=None, dyn_reta_end=None, **kwargs):
super().__init__(**kwargs)
self.reta = reta
self.dyn_reta_start = dyn_reta_start
self.dyn_reta_end = dyn_reta_end
def get_dyn_reta(self, ss):
return self.reta * self.get_dyn_value(
ss, self.dyn_reta_start, self.dyn_reta_end
)
class EulerStep(SingleStepSampler):
name = "euler"
step = SingleStepSampler.euler_step
class DPMPPStepBase(SingleStepSampler):
@staticmethod
def sigma_fn(t):
return t.neg().exp()
@staticmethod
def t_fn(t):
return t.log().neg()
class DPMPP2MStep(DPMPPStepBase):
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
t, t_next = self.t_fn(ss.sigma), self.t_fn(ss.sigma_next)
h = t_next - t
st, st_next = self.sigma_fn(t), self.sigma_fn(t_next)
if len(ss.dhist) == 0 or ss.sigma_prev is None:
return (st_next / st) * x - (-h).expm1() * ss.denoised, 0.0
h_last = t - self.t_fn(ss.sigma_prev)
r = h_last / h
denoised, old_denoised = ss.denoised, ss.dhist[-1]
denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised
return (st_next / st) * x - (-h).expm1() * denoised_d, 0.0
class DPMPP2MSDEStep(SingleStepSampler):
name = "dpmpp_2m_sde"
def __init__(self, *, solver_type="midpoint", **kwargs):
super().__init__(**kwargs)
self.solver_type = solver_type
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
denoised = ss.denoised
if ss.sigma_next == 0:
return denoised, None
# DPM-Solver++(2M) SDE
t, s = -ss.sigma.log(), -ss.sigma_next.log()
h = s - t
eta_h = self.get_dyn_eta(ss) * h
x = (
ss.sigma_next / ss.sigma * (-eta_h).exp() * x
+ (-h - eta_h).expm1().neg() * denoised
)
noise_strength = ss.sigma_next * (-2 * eta_h).expm1().neg().sqrt()
if len(ss.dhist) == 0 or ss.sigma_prev is None:
return x, noise_strength
h_last = (-ss.sigma.log()) - (-ss.sigma_prev.log())
r = h_last / h
old_denoised = ss.dhist[-1]
if self.solver_type == "heun":
x = x + (
((-h - eta_h).expm1().neg() / (-h - eta_h) + 1)
* (1 / r)
* (denoised - old_denoised)
)
elif self.solver_type == "midpoint":
x = x + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * (
denoised - old_denoised
)
return x, noise_strength
class DPMPP3MSDEStep(SingleStepSampler):
name = "dpmpp_3m_sde"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
denoised = ss.denoised
if ss.sigma_next == 0:
return denoised, 0
t, s = -ss.sigma.log(), -ss.sigma_next.log()
h = s - t
eta = self.get_dyn_eta(ss)
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()
if len(ss.dhist) == 0 or ss.sigma_prev is None:
return x, noise_strength
h_1 = (-ss.sigma.log()) - (-ss.sigma_prev.log())
denoised_1 = ss.dhist[-1]
if len(ss.dhist) == 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:
h_2 = (-ss.sigma_prev.log()) - (-ss.sigmas[ss.idx - 2].log())
denoised_2 = ss.dhist[-2]
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
return x, noise_strength
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class ReversibleHeunStep(ReversibleSingleStepSampler):
name = "reversible_heun"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(
self.get_dyn_reta(ss)
)
dt = sigma_down - ss.sigma
dt_reversible = sigma_down_reversible - ss.sigma
# Calculate the derivative using the model
d = to_d(x, ss.sigma, ss.denoised)
# Predict the sample at the next sigma using Euler step
x_pred = x + d * dt
# Denoised sample at the next sigma
denoised_next = ss.model(x_pred, sigma_down, model_call_idx=1)
# Calculate the derivative at the next sigma
d_next = to_d(x_pred, sigma_down, denoised_next)
# Update the sample using the Reversible Heun formula
x = x + dt * (d + d_next) / 2 - dt_reversible**2 * (d_next - d) / 4
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class ReversibleHeun1SStep(ReversibleSingleStepSampler):
name = "reversible_heun_1s"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
# Reversible Heun-inspired update (first-order)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(
self.get_dyn_reta(ss)
)
sigma_i, sigma_i_plus_1 = ss.sigma, sigma_down
dt = sigma_i_plus_1 - sigma_i
dt_reversible = sigma_down_reversible - sigma_i
eff_x = ss.xhist[-1] if len(ss.xhist) else x
# Calculate the derivative using the model
d_i_old = to_d(
eff_x,
sigma_i,
ss.dhist[-1]
if len(ss.dhist)
else ss.model(eff_x, sigma_i, model_call_idx=1),
)
# Predict the sample at the next sigma using Euler step
x_pred = eff_x + d_i_old * dt
# Calculate the derivative at the next sigma
d_i_plus_1 = to_d(x_pred, sigma_i_plus_1, ss.denoised)
# Update the sample using the Reversible Heun formula
x = (
x
+ dt * (d_i_old + d_i_plus_1) / 2
- dt_reversible**2 * (d_i_plus_1 - d_i_old) / 4
)
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class RESStep(SingleStepSampler):
name = "res"
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
pass
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
eta = self.get_dyn_eta(ss)
sigma_down, sigma_up = ss.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 = _de_second_order(
h=h, c2=self.c2, simple_phi_calc=self.simple_phi
)
c2_h = 0.5 * h
x_2 = math.exp(-c2_h) * x + a2_1 * h * denoised
lam_2 = lam + c2_h
sigma_2 = lam_2.neg().exp()
denoised2 = ss.model(x_2, sigma_2, model_call_idx=1)
x = math.exp(-h) * x + h * (b1 * denoised + b2 * denoised2)
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class TrapezoidalStep(SingleStepSampler):
name = "trapezoidal"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
dt = ss.sigma_next - ss.sigma
denoised = ss.denoised
# Calculate the derivative using the model
d_i = to_d(x, ss.sigma, denoised)
# Predict the sample at the next sigma using Euler step
x_pred = x + d_i * dt
# Denoised sample at the next sigma
denoised_next = ss.model(x_pred, ss.sigma_next, model_call_idx=1)
# Calculate the derivative at the next sigma
d_next = to_d(x_pred, ss.sigma_next, denoised_next)
dt_2 = sigma_down - ss.sigma
# Update the sample using the Trapezoidal rule
x = x + dt_2 * (d_i + d_next) / 2
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class BogackiStep(ReversibleSingleStepSampler):
name = "bogacki"
reversible = False
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(
self.get_dyn_reta(ss)
)
sigma, sigma_next = ss.sigma, sigma_down
dt = sigma_next - sigma
dt_reversible = sigma_down_reversible - sigma
denoised = ss.denoised
# Calculate the derivative using the model
d = to_d(x, sigma, denoised)
# Bogacki-Shampine steps
k1 = d * dt
k2 = (
to_d(
x + k1 / 2,
sigma + dt / 2,
ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=1),
)
* dt
)
k3 = (
to_d(
x + 3 * k1 / 4 + k2 / 4,
sigma + 3 * dt / 4,
ss.model(x + 3 * k1 / 4 + k2 / 4, sigma + 3 * dt / 4, model_call_idx=2),
)
* dt
)
# Reversible correction term (inspired by Reversible Heun)
correction = dt_reversible**2 * (k3 - k2) / 6 if self.reversible else 0.0
# Update the sample
x = x + 2 * k1 / 9 + k2 / 3 + 4 * k3 / 9 - correction
return x, sigma_up
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"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
sigma = ss.sigma
# Calculate the derivative using the model
d = to_d(x, sigma, ss.denoised)
dt = sigma_down - sigma
# Runge-Kutta steps
k1 = d * dt
k2 = (
to_d(
x + k1 / 2,
sigma + dt / 2,
ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=1),
)
* dt
)
k3 = (
to_d(
x + k2 / 2,
sigma + dt / 2,
ss.model(x + k2 / 2, sigma + dt / 2, model_call_idx=2),
)
* dt
)
k4 = (
to_d(
x + k3,
sigma + dt,
ss.model(x + k3, sigma + dt, model_call_idx=3),
)
* dt
)
# Update the sample
x = x + (k1 + 2 * k2 + 2 * k3 + k4) / 6
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class EulerDancingStep(SingleStepSampler):
name = "euler_dancing"
def __init__(
self,
*,
deta=1.0,
ds_noise=1.0,
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
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):
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:
return x, 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
x = x + self.noise_sampler(ss.sigma, sigma_leap) * 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
class DPMPP2SStep(DPMPPStepBase):
name = "dpmpp_2s"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
t_fn, sigma_fn = self.t_fn, self.sigma_fn
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
# 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
x_2 = (sigma_fn(s) / sigma_fn(t)) * x - (-h * r).expm1() * ss.denoised
denoised_2 = ss.model(x_2, sigma_fn(s), model_call_idx=0)
x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_2
return x, sigma_up
class DPMPPSDEStep(DPMPPStepBase):
name = "dpmpp_sde"
def __init__(self, *args, r=1 / 2, **kwargs):
super().__init__(*args, **kwargs)
self.r = r
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
t_fn, sigma_fn = self.t_fn, self.sigma_fn
r, eta, s_noise = self.r, self.get_dyn_eta(ss), self.s_noise
noise_sampler = self.noise_sampler
sigma_down, sigma_up = ss.get_ancestral_step(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)
x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - (t - s_).expm1() * ss.denoised
x_2 = x_2 + noise_sampler(sigma_fn(t), sigma_fn(s)) * s_noise * su
denoised_2 = ss.model(x_2, sigma_fn(s), model_call_idx=1)
# 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)) * x - (t - t_next_).expm1() * denoised_d
return x, su
# Based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
# Which was originally written by Katherine Crowson
class TTMJVPStep(SingleStepSampler):
name = "ttm_jvp"
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):
if ss.sigma_next == 0:
return ss.denoised, ss.sigma.new_zeros(1)
eta = self.get_dyn_eta(ss)
sigma_down, sigma_up = ss.get_ancestral_step(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, denoised_prime = ss.model(
x, sigma, tangents=(eps * -sigma, -sigma), model_call_idx=1
)
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
if not eta:
return x, ss.sigma.new_zeros(1)
phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * eta))
return x, sigma_next * phi_1_noise
STEP_SAMPLERS = {
"euler": EulerStep,
"dpmpp_sde": DPMPPSDEStep,
"dpmpp_2m": DPMPP2MStep,
"dpmpp_2m_sde": DPMPP2MSDEStep,
"dpmpp_3m_sde": DPMPP3MSDEStep,
"dpmpp_2s": DPMPP2SStep,
"reversible_heun": ReversibleHeunStep,
"reversible_heun_1s": ReversibleHeun1SStep,
"res": RESStep,
"trapezoidal": TrapezoidalStep,
"bogacki": BogackiStep,
"reversible_bogacki": ReversibleBogackiStep,
"rk4": RK4Step,
"euler_dancing": EulerDancingStep,
"ttm_jvp": TTMJVPStep,
}
__all__ = (
"STEP_SAMPLERS",
"EulerStep",
"DPMPP2MStep",
"DPMPP2MSDEStep",
"DPMPP3MSDEStep",
"DPMPP2SStep",
"ReversibleHeunStep",
"ReversibleHeun1SStep",
"RESStep",
"TrapezoidalStep",
"BogackiStep",
"ReversibleBogackiStep",
"EulerDancingStep",
"TTMJVPStep",
)
+141 -122
View File
@@ -2,188 +2,207 @@ import torch
from comfy.k_diffusion.sampling import get_ancestral_step
from .filtering import FilterRefs
from .model import History
class StepSamplerChain:
class Items:
def __init__(self, items=None):
self.items = [] if items is None else items
def clone(self):
return self.__class__(items=self.items.copy())
def append(self, item):
self.items.append(item)
return item
class History:
def __init__(self, x, size):
self.history = torch.zeros(size, *x.shape, device=x.device, dtype=x.dtype)
self.size = size
self.pos = 0
self.last = None
def __getitem__(self, key):
return self.items[key]
def __setitem__(self, key, value):
self.items[key] = value
def __len__(self):
return min(self.pos, self.size)
return len(self.items)
def __getitem__(self, k):
idx = (self.pos + k if k < 0 else self.pos + -self.size + k) % self.size
# print(f"\nFETCH {k}: pos={self.pos}, size={self.size}, at={idx}")
return self.history[idx]
def push(self, val):
# print(f"\nPUSH {self.pos % self.size}: pos={self.pos}, size={self.size}")
self.last = self.pos % self.size
self.history[self.last] = val
self.pos += 1
def reset(self):
self.pos = 0
self.last = None
def __iter__(self):
return self.items.__iter__()
class ModelCallCache:
class CommonOptionsItems(Items):
def __init__(self, *, s_noise=1.0, eta=1.0, items=None, **kwargs):
super().__init__(items=items)
self.options = kwargs
self.s_noise = s_noise
self.eta = eta
def clone(self):
obj = super().clone()
obj.options = self.options.copy()
obj.s_noise = self.s_noise
obj.eta = self.eta
return obj
class StepSamplerChain(CommonOptionsItems):
def __init__(
self, model, x, s_in, extra_args, *, size=0, max_use=1000000, threshold=1
self,
*,
merge_method="divide",
time_mode="step",
time_start=0,
time_end=999,
**kwargs,
):
self.size = size
self.model = model
self.threshold = threshold
self.s_in = s_in
self.extra_args = extra_args
self.max_use = max_use
if self.size < 1:
return
self.mcc = torch.zeros(size, *x.shape, device=x.device, dtype=x.dtype)
self.jmcc = torch.zeros_like(self.mcc)
self.reset_cache()
super().__init__(**kwargs)
self.merge_method = merge_method
if time_mode not in ("step", "step_pct", "sigma"):
raise ValueError("Bad time mode")
self.time_mode = time_mode
self.time_start, self.time_end = time_start, time_end
def reset_cache(self):
size = self.size
self.slot = [None] * size
self.jslot = [None] * size
self.slot_use = [self.max_use] * size
def clone(self):
obj = super().clone()
obj.merge_method = self.merge_method
obj.time_mode = self.time_mode
obj.time_start, obj.time_end = self.time_start, self.time_end
obj.options = self.options.copy()
return obj
def get(self, idx, *, jvp=False):
idx -= self.threshold
if (
idx >= self.size
or idx < 0
or self.slot[idx] is None
or self.slot_use[idx] < 1
):
return None
if jvp and self.jslot[idx] is None:
return None
self.slot_use[idx] -= 1
return self.slot[idx] if not jvp else (self.slot[idx], self.jslot[idx])
def set(self, idx, denoised, jdenoised=None):
idx -= self.threshold
if idx < 0 or idx >= self.size:
return
self.slot_use[idx] = self.max_use
self.slot[idx] = denoised
self.jslot[idx] = jdenoised
class ParamGroup(Items):
pass
def call_model(self, x, sigma, **kwargs):
return self.model(x, sigma * self.s_in, **self.extra_args, **kwargs)
def __call__(self, x, sigma, *, model_call_idx=0, tangents=None, **kwargs):
result = self.get(model_call_idx, jvp=tangents is not None)
# print(
# f"MODEL: idx={model_call_idx}, size={self.size}, threshold={self.threshold}, cached={result is not None}"
# )
if result is not None:
return result
if tangents is None:
denoised = self.call_model(x, sigma, **kwargs)
self.set(model_call_idx, denoised)
return denoised
denoised, denoised_prime = torch.func.jvp(self.call_model, (x, sigma), tangents)
self.set(model_call_idx, denoised, jdenoised=denoised_prime)
return denoised, denoised_prime
class StepSamplerGroups(CommonOptionsItems):
pass
class SamplerState:
CLONE_KEYS = (
"model",
"hist",
"extra_args",
"disable_status",
"eta",
"reta",
"s_noise",
"sigmas",
"callback_",
"noise_sampler",
"noise",
"idx",
"total_steps",
"step",
"substep",
"sigma",
"sigma_next",
"sigma_prev",
"sigma_down",
"sigma_up",
"refs",
)
def __init__(
self,
model,
sigmas,
idx,
dhist,
xhist,
extra_args,
*,
step=0,
substep=0,
noise_sampler,
callback=None,
denoised=None,
noise=None,
eta=1.0,
reta=1.0,
s_noise=1.0,
disable_status=False,
history_size=4,
):
self.model = model
self.dhist = dhist
self.xhist = xhist
self.hist = History(max(1, history_size))
self.extra_args = extra_args
self.eta = eta
self.reta = reta
self.s_noise = s_noise
self.sigmas = sigmas
self.denoised = denoised
self.callback_ = callback
self.noise_sampler = noise_sampler
self.update(idx)
self.noise = noise
self.disable_status = disable_status
self.step = 0
self.substep = 0
self.total_steps = len(sigmas) - 1
self.update(idx) # Sets idx, sigma_prev, sigma, sigma_down, refs
def update(self, idx=None):
@property
def hcur(self):
return self.hist[-1]
@property
def hprev(self):
return self.hist[-2]
@property
def denoised(self):
return self.hcur.denoised
@property
def dt(self):
return self.sigma_next - self.sigma
@property
def d(self):
return self.hcur.d
def update(self, idx=None, step=None, substep=None):
idx = self.idx if idx is None else idx
self.idx = idx
self.sigma_prev = None if idx < 1 else self.sigmas[idx - 1]
self.sigma, self.sigma_next = self.sigmas[idx], self.sigmas[idx + 1]
# if self.sigma_prev is not None and self.sigma < self.sigma_prev:
# self.dhist.reset()
# self.xhist.reset()
self.sigma_down, self.sigma_up = get_ancestral_step(
self.sigma, self.sigma_next, eta=self.eta
)
self.sigma_down_reversible, self.sigma_up_reversible = get_ancestral_step(
self.sigma, self.sigma_next, eta=self.reta
)
if step is not None:
self.step = step
if substep is not None:
self.substep = substep
self.refs = FilterRefs.from_ss(self)
def get_ancestral_step(self, eta=1.0):
return get_ancestral_step(self.sigma, self.sigma_next, eta=eta)
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
)
)
return sd, su
def clone_edit(self, **kwargs):
obj = self.__class__.__new__(self.__class__)
for k in (
"model",
"dhist",
"xhist",
"extra_args",
"eta",
"reta",
"s_noise",
"sigmas",
"denoised",
"callback_",
"noise_sampler",
"idx",
"sigma",
"sigma_next",
"sigma_prev",
"sigma_down",
"sigma_up",
"sigma_down_reversible",
"sigma_up_reversible",
):
for k in self.CLONE_KEYS:
setattr(obj, k, kwargs[k] if k in kwargs else getattr(self, k))
obj.update()
return obj
def callback(self, x):
def callback(self, hi=None):
if not self.callback_:
return None
return self.callback_(
{
"x": x,
"i": self.idx,
"sigma": self.sigma,
"sigma_hat": self.sigma,
"denoised": self.dhist[-1],
}
)
hi = self.hcur if hi is None else hi
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
+67 -9
View File
@@ -1,17 +1,24 @@
import math
import contextlib
import torch
from comfy.k_diffusion.sampling import to_d
def scale_noise(noise, factor=1.0, *, normalized=True, threshold_std_devs=2.5):
def scale_noise(
noise,
factor=1.0,
*,
normalized=True,
normalize_dims=(-3, -2, -1),
):
if not normalized or noise.numel() == 0:
return noise.mul_(factor) if factor != 1 else noise
mean, std = noise.mean().item(), noise.std().item()
threshold = threshold_std_devs / math.sqrt(noise.numel())
if abs(mean) > threshold:
noise -= mean
if abs(1.0 - std) > threshold:
noise /= std
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):
@@ -20,3 +27,54 @@ def find_first_unsorted(tensor, desc=True):
fun = torch.gt if desc else torch.lt
first_unsorted = fun(tensor[1:], tensor[:-1]).nonzero().flatten()[:1].add_(1)
return None if not len(first_unsorted) else first_unsorted.item()
def fallback(val, default, exclude=None):
return val if val is not exclude else default
def step_generator(gen, *, get_next, initial=None):
next_val = initial
with contextlib.suppress(StopIteration):
while True:
result = gen.send(next_val)
next_val = get_next(result)
yield result
# From Gaeros. Thanks!
def extract_pred(x_before, x_after, sigma_before, sigma_after):
if sigma_after == 0:
return x_after, torch.zeros_like(x_after)
alpha = sigma_after / sigma_before
denoised = (x_after - alpha * x_before) / (1 - alpha)
return denoised, to_d(x_after, sigma_after, denoised)
def resolve_value(keys, obj):
if not len(keys):
raise ValueError("Cannot resolve empty key list")
result = obj
class Empty:
pass
for idx, key in enumerate(keys):
if not (hasattr(result, "__getattr__") or hasattr(obj, "__getattribute__")):
raise ValueError(
f"Cannot access key {key}: value does not support attribute access"
)
result = getattr(result, key, Empty)
if result is Empty:
raise AttributeError(f"Key {key} from path {'.'.join(keys)} does not exist")
def check_time(time_mode, time_start, time_end, sigma, step, steps):
step_pct = step / steps if steps != 0 else 0.0
if time_mode == "step":
return time_start <= step <= time_end
if time_mode == "step_pct":
return time_start <= step_pct <= time_end
if time_mode == "sigma":
return time_start >= sigma >= time_end
raise ValueError("Bad time mode")