@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||

|
||||
|
||||
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:
|
||||
|
||||

|
||||
|
||||
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
@@ -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 |
@@ -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
@@ -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
|
||||
```
|
||||
@@ -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",
|
||||
)
|
||||
@@ -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, (")", "]"))
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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",
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
+460
-196
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user