Author SHA1 Message Date
blepping d8aadb2359 Perlin documentation and cleanups
Add OCSNoise PerlinSimple node
2024-08-27 07:52:12 -06:00
blepping 9b7e6b6083 Enable alt_cfgpp_scale for dpmpp_2s, dpmpp_sde and res samplers 2024-08-27 04:59:06 -06:00
blepping 7fc26d0488 Euler should allow alt_cfgpp 2024-08-25 20:27:00 -06:00
blepping 7418974f6b Initial Perlin3D and 2D implementation 2024-08-21 06:50:54 -06:00
blepping bc1e4b0874 Fix setting group parameters 2024-08-17 14:45:39 -06:00
blepping 3e73d2e352 Fix uncond preview_mode 2024-08-17 08:20:55 -06:00
blepping 0aa45f98ff Allow setting a preview_mode in groups 2024-08-17 07:45:19 -06:00
blepping aac7a05bb2 Fix issue where batches didn't work with Diffrax solver
Fix some other issues related to ODE/SDE solvers and batches
2024-08-17 07:23:51 -06:00
blepping ac1951101f Add TAESD encode/decode and some image handling function for expressions
Documentation updates

Allow filtering model input result
2024-08-16 11:30:04 -06:00
blepping fdbbd970e7 Adjust noise scaling approach 2024-08-16 11:09:54 -06:00
blepping 9db4fbe514 Fix noise sampler reset for really reals this time 2024-08-16 11:09:19 -06:00
blepping 2d8125e011 Allow dash in expression symbols 2024-08-16 11:07:58 -06:00
blepping 076d80918a Add tooltip metadata to nodes 2024-08-15 08:40:15 -06:00
blepping bf43d36bac Allow using nnlatentupscale in expressions if available 2024-08-15 07:53:46 -06:00
blepping fd5d0ef6a1 Fix noise cache update to work better when the latent size changes 2024-08-15 05:01:13 -06:00
blepping 0595a6d304 Fix another fstr in fstr issue 2024-08-09 08:52:51 -06:00
blepping 305720c66e Fix issue with fstr in fstr 2024-08-09 08:33:51 -06:00
blepping f0e52ec054 Make restarting work better when the latent changes sizes during sampling 2024-08-08 19:35:51 -06:00
blepping 442bde294a Add t_shape expression function
Fix width/height order for t_scale in absolute mode
2024-08-08 14:10:25 -06:00
blepping 764f04a6a7 Fix t_noise expression function 2024-08-08 07:52:40 -06:00
blepping 2481e7263b Allow setting temporary variables in expressions
Improve sequence operator in expressions

Documentation improvements

Turn down eval debug spam a bit

Add ternary operator

Support static eval for constants in some more operator types
2024-08-08 06:27:53 -06:00
blepping 470b38231f Refactor (#1)
Refactor all the things!
2024-08-04 11:28:06 -06:00
32 changed files with 8919 additions and 1161 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
MIT License
Copyright (c) 2024 blepping
Copyright (c) 2024 blepping <https://github.com/blepping>
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
+662 -77
View File
@@ -1,93 +1,678 @@
# 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.
Feel free create a question in Discussions for usage help: [OCS Q&A Discussion](https://github.com/blepping/comfyui_overly_complicated_sampling/discussions/categories/q-a)
*Note*: You will basically always have to tweak settings like `s_noise` to get a good result. If the generation looks smooth/undetailed increase `s_noise` somewhere. If it looks crunchy, super high contrast, etc then try reducing noise.
## Features
## Nodes
* Many different samplers.
* Allows scheduling samplers (i.e. run `euler` for steps 1-4, then switch to `dpmpp_sde`).
* CFG++ support (for some samplers).
* Native support for Restart sigmas.
* Supports custom noise types.
* Immiscible noise for sampling and Restart. See https://arxiv.org/abs/2406.12303 (note that it was designed for training not inference).
* Allows splitting/combining steps in various ways for (potentially) more accurate sampling.
* Supports Diffrax, torchdiffeq, torchode and torchsde solver backends. (SDE mode not recommended currently.)
* Many tuneable parameters to play with.
### ComposableSampler
**Possible Parameters**
* `avgmerge_stretch`(`0.4`): Used for `average` and `sample` merge types. See below.
* `model_call_cache`(unset): Caches the result of model calls. For example, Bogacki is 3 model calls per step. The first one usually depends on the merge strategy: `average` for example shares the first model evaluation between substeps, but subsequent model calls (i.e. Bogacki 2nd and 3rd model evaluations) still occur. When the model call cache is active, it's possible to cache those evalutions and avoid a model call for the remaining substeps. If you set `model_call_cache` to `1` then the result of that second call will be cached and if you're running two Bogacki substeps then the second one will use the cached version. Massively accelerates inference (especially when using the `average` merge strategy) but is likely very unsound and inaccurate. Does not apply to the sampler call for the `sample` merge strategy.
* `model_call_cache_threshold`(`1`): Disables caching model call results with a call index below the threshold value (starting at 0). For example, if set to `2` and using a sampler like Bogacki that calls the model two extra times, the first will never be cached. The default value of `1` disabling caching for the first model call per substep. I generally would not recommend setting it to `0`, especially with `average` or `sample` merge strategies.
* `model_call_cache_max_use`(`1000000`): The number of times cache items can be re-used. The default is effectively no limit. Where would this be useful? Let's say you're using the `average` merge strategy and a multi step sampler that calls the model at least one more time with 50 substeps. If you set the value to `25`, the model cache result will be updated around substep 25 which _may_ produce better results than reusing the result 50 times.
Since it's kind of confusing even for me, a little more explanation: The model call cache caches results for model call indexes between `model_call_cache_threshold` and `model_call_cache - 1`. If you set `model_call_cache_threshold` to `0` and `model_call_cache` to `1` then only the first model call will be cached. If you set `model_call_cache_threshold` to `1` and `model_call_cache` to `2` then call 0 will not be cached, call 1 will be cached, call 2 will be cached, call 3 will not be cached, and so on.
#### Merging
When running multiple substeps per step, the results will combined based on the merge strategy. Possible strategies (in order of least weird to most weird):
* `divide`: Creates a linear schedule between the current sigma and the next and runs the substeps in sequence. The model is called at least once per substep.
* `normal`: The model is called at least once per substep (and possibly additional times for higher order samplers). The result of each substep is noised and the next substep uses that result. Then all the results are averaged.
* `average`: The model is called once at the beginning of the step and substeps share that result (but it may be called additional times for higher order samplers). This means substeps for samplers like reversible Euler, Heun 1s, DPM++ 2m SDE are essentially free. May be theoretically very unsound and inaccurate, requires manual tweaking of settings like `s_noise`. Supports the parameter `avgmerge_stretch`(`0.4`) which basically rolls back the current sigma and adds some noise (otherwise running a substep is deterministic and there would be no point to running a sampler like Euler more than once).
* `sample`: Like `average` (and uses `avgmerge_stretch`) but instead of simply using the average, it does a sampler step toward that instead. You can plug in any substep sampler to the `merge_sampler_opt` input (if unconnected and the merge method is `sample` then Euler will be used). *Note*: Substeps in the attached sampler will be ignored.
* `sample_uncached`: Similar to `sample`, however it calls the model per substep instead of caching the result and sharing it. Aside from sampling toward the result, it works more like the `normal` merge strategy. Theoretically it should be better because it's taking less shortcuts but results seem worse.
When using `average` and `sample` merge strategies and with model call caching enabled you can get away with setting substeps super high. Running something like 100 substeps is actually quite practical and seems to work well.
### ComposableStepSampler
This node has a text input for YAML (or JSON) advanced parameters.
For example, you could enter something like this in the field:
```yaml
reta: 1.1
leap: 3
dyn_deta_mode: "deta"
```
**Possible Parameters**
#### General
* `eta`(`1.0`): Will override `eta` in the node if set.
* `dyn_eta_start`(`unset`) and `dyn_eta_end`(`unset`): No effect unless both values are set. Will interpolate between start and end based on the percentage of sampling. *Note*: This is a factor applied to ETA, not a flat value.
* `s_noise`(`1.0`): Will override `s_noise` in the node if set.
* `solver_type`(`midpoint`): Applies to DPM++ 2m SDE. May be one of `midpoint` or `heun` (`midpoint` is generally recommended).
#### Reversible
* `reta`(`1.0`): Reverse ETA.
* `dyn_reta_start`(`unset`) and `dyn_reta_end`(`unset`): No effect unless both values are set. Will interpolate between start and end based on the percentage of sampling. *Note*: This is a factor applied to RETA, not a flat value.
#### Dancing
* `leap`(`2`): Distance to try to leap forward. If you set `leap` to `1` you just get plain old Euler ancestral.
* `deta`(`1.0`): ETA used for dance steps.
* `dyn_deta_start`(`unset`) and `dyn_deta_end`(`unset`): No effect unless both values are set. Will interpolate between start and end based on the percentage of sampling. *Note*: This is a factor applied to DETA, not a flat value.
* `dyn_deta_mode`(`lerp`): May be one of:
* `deta`: Scales `deta` based on the value from `dyn_deta_start/end`.
* `lerp`: Does the dance step according to `deta` and then LERPs the non-dance sample result with the dance sample result based on the scale calculated from `dyn_deta_start/end` (which is `1.0` if they are unset). For example, if the dance scale is `0.5` you will get 50% normal sampling, 50% dancing sampling.
* `lerp_alt`: Similar to `lerp` except it LERPs with the leap result instead of a normal Euler ancestral result.
#### RES
* `res_simple_phi`(`false`): Uses a faster but possibly less accurate method for calculating phi. What does phi do? I haven't the foggiest!
* `res_c2`(`0.5`): Solver partial step size, the default of `0.5` appears to use the midpoint. Setting it to a lower value might possibly be more accurate but slower?
#### TTM JVP
`alterate_phi_2_calc`(`true`): Supposedly works better than disabled when ETA isn't 0. I didn't notice a difference.
**Note**: TTM is a weird sampler. If you're using model caching you must make sure the entries TTM uses are populated first (by having before any other samplers that call the model multiple times). It may also not work with some other model patches and upscale methods.
## Credits
I can move code around but sampling math and creating samplers is far beyond my ability. I didn't write any of the original samplers:
* Euler, DPMPP SDE, DPMPP 2S, DPM++ 2m, 2m SDE and 3m SDE samplers based on ComfyUI's implementation.
* Euler, Heun++2, DPMPP SDE, DPMPP 2S, DPM++ 2m, 2m SDE and 3m SDE samplers based on ComfyUI's implementation.
* Reversible Heun, Reversible Heun 1s, RES, Trapezoidal, Bogacki, Reversible Bogacki, RK4 and Euler Dancing samplers based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
* TTM JVP sampler based on implementation written by Katherine Crowson (but yoinked from the Extra-Samplers repo mentioned above).
* IPNDM, IPNDM_V and DEIS adapted from https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py (I used the Comfy version as a reference).
* Normal substep merge strategy based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
* Immiscible noise processing based on implementation from https://github.com/kohya-ss/sd-scripts/pull/1395 and idea for sampling with it from https://github.com/Clybius
* Precedence climbing (Pratt) expression parser based on implementation from https://github.com/andychu/pratt-parsing-demo
Thanks!
This repo wouldn't be possible without building on the work of others. Thanks!
## Usage
First, a note on the basic structure:
![Basic nodes example](assets/basic_sampling.png)
The sampler node connects to a group node. You can chain group nodes, however only one can match per step. Checking for a match starts at the group furthest from the sampler node. So if you have `Group1 -> Group2 -> Group3 -> Sampler`, the groups will be tried starting from `Group1`. Currently matching groups is only time based.
You will then connect a substeps node to the group. These can also be chained and like groups, execution starts with the node furthest from the group. I.E.: `Substeps1 -> Substeps2 -> Substeps3 -> Group` will start with `Substeps1`.
Most of these nodes have a text parameter input (YAML format - JSON is also valid YAML so you can use that instead if you prefer) and a parameter input. The parameter input can be used to specify stuff like custom noise types.
You may use filters and expressions in the text parameter input. See:
* [Filters](docs/filter.md)
* [Expressions](docs/expression.md)
## Integration
* [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) - allows access to many more blend and scaling modes as well as some extra features.
* [ComfyUI-sonar](https://github.com/blepping/ComfyUI-sonar) - allows access to many more noise types as well as the Power Filter feature.
* [ComfyUi_NNLatentUpscale](https://github.com/Ttl/ComfyUi_NNLatentUpscale) - allows access to the `t_scale_nnlatentupscale` function in expressions.
If you're going to use OCS, I strongly recommend also installing `ComfyUI-bleh` and `ComfyUI-sonar` as they increase the functionality a lot.
## Nodes
### `OCS Sampler`
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: false
# Interval (in full steps) to reset the cache. Brownian noise takes time into account so
# if using Brownian with caching enabled you will generally want to reset each step.
cache_reset_interval: 9999
# Immiscible noise processing, see: https://arxiv.org/abs/2406.12303
immiscible:
# Batch size, 0 disables.
size: 0
# Reference mode, values can be one of:
# x: Uses the current latent as a reference.
# noise: Uses the current noise as a reference (x - denoised)
# denoised: Uses the model image prediction as a reference (factors in positive and negative prompts).
# uncond: The model unconditional prediction (negative prompt)
# cond: The model conditional prediction (positive prompt)
# Advanced feature: Additionally you may enter a string of operations in the format:
# "x - denoised * 2 + cond" (just an example, not a recommended setting)
# Possible operations: + - / * min max add sub div mul
# Note: Each value and operation must be space delimited (i.e. "x-1" will not work).
# Also normal operator precedence does not apply here.
ref: default
# Batching mode, one of:
# batch: Matches vs batches. Immiscible mode is disabled if size < 2
# channel: Splits the batch into a list of channels and matches against those.
# row: Splits the batch into a list of rows and matches against those.
# column: Splits the batch into a list of columns and matches against those.
# Note: Requires reshaping both the noise and x, may be slow and consume
# a lot of VRAM.
batching: channel
# Scale for reference latent. Can be negative.
scale_ref: 1.0
# Allows normalizing the reference. If this is a list, you can specify the dimensions to
# normalize. See normalize_dims above and https://pytorch.org/docs/stable/generated/torch.std.html
normalize_ref: false
# The proportion of immiscible-ized noise.
# You get (immiscible_noise * strength) + ((1.0 - strength) * normal_noise) - LERP.
strength: 1.0
# See: https://docs.scipy.org/doc/scipy/reference/generated/scipy.optimize.linear_sum_assignment.html#scipy.optimize.linear_sum_assignment
maximize: false
filter: null
# Model calls can be cached. This is very experimental: I don't recommend using it
# unless you know what you're doing.
model:
cache:
# The cache size.
size: 0
# Threshold for model call caching. For example if you have size=3 and threshold=1
# then model calls 1 through 3 will be cached, but model call 0 will not be (the first one).
# Additional explanation: Some samplers call the model multiple times per step. For example,
# Bogacki uses three model calls: 0, 1, 2
threshold: 1
# Maximum use count for cache items.
max_use: 1000000
filter:
input: null
denoised: null
jdenoised: null
cond: null
uncond: null
```
</details><br/>
Any parameters you don't specify will use the defaults. For example if your text parameter block is:
```yaml
noise:
cpu_noise: false
```
Then the rest of the parameters will use the defaults shown above.
***
### `OCS Group`
Defines a group of substeps.
#### Merging
When running multiple substeps per step, the results will combined based on the merge strategy. Possible strategies (in order of least weird to most weird):
* `simple`: Doesn't merge anything: only runs a single substep per step.
* `divide`: Creates a linear schedule between the current sigma and the next and runs the substeps in sequence. The model is called at least once per substep.
* `normal`: The model is called at least once per step (and possibly additional times for higher order samplers). Each substep shares the first model call result. The results are averaged together. *Note*: Since the first model call is shared and the initial input is the same for each substep, there is no point in running multiple identical substeps. Also note: This merge strategy doesn't work well with non-ancestral samplers (i.e. dpmpp_2m or any sampler with `eta: 0`).
* `overshoot`: The model is called at least once per step. It will sample steps equal to the number of substeps, starting from the current step. Then it will restart back to the expected step.
<!--
* `average`: The model is called once at the beginning of the step and substeps share that result (but it may be called additional times for higher order samplers). This means substeps for samplers like reversible Euler, Heun 1s, DPM++ 2m SDE are essentially free. May be theoretically very unsound and inaccurate, requires manual tweaking of settings like `s_noise`. Supports the parameter `avgmerge_stretch`(`0.4`) which basically rolls back the current sigma and adds some noise (otherwise running a substep is deterministic and there would be no point to running a sampler like Euler more than once).
* `sample`: Like `average` (and uses `avgmerge_stretch`) but instead of simply using the average, it does a sampler step toward that instead. You can plug in any substep sampler to the `merge_sampler_opt` input (if unconnected and the merge method is `sample` then Euler will be used). *Note*: Substeps in the attached sampler will be ignored.
* `sample_uncached`: Similar to `sample`, however it calls the model per substep instead of caching the result and sharing it. Aside from sampling toward the result, it works more like the `normal` merge strategy. Theoretically it should be better because it's taking less shortcuts but results seem worse.
When using `average` and `sample` merge strategies and with model call caching enabled you can get away with setting substeps super high. Running something like 100 substeps is actually quite practical and seems to work well.
-->
#### Node Parameters
* `merge_method`: One of the merge methods described above in the Merging section.
* `time_mode`(`step`): One of `step`, `step_pct`, `sigma`. Time matching mode. Matching based on steps generally will be simplest. Matches are inclusive and steps start at 0 (so step 0 is the first step). `step_pct` is the percentage of total steps (1.0=100%, 0.5=50%, etc).
* `time_start`(`0`): Match start time.
* `time_end`(`999`): Match end time.
Example:
![Group time filter example](assets/group_time_example.png)
The left side group matches steps 0, 1, 2. The right side group matches all steps. This setup will use whatever substeps are connected to the first group for the first three steps and the second group will handle the rest.
#### Input Parameters
<!--
* `merge_sampler`: Value type: `OCS_SUBSTEPS`. Only used when `merge_method` is `sample` or `sample_uncached`. Allows defining the sampler used for merging substeps.
-->
* `restart_custom_noise`: Currently only used by the `overshoot` merge method.
#### Text Parameters
Shown in YAML with default values.
<details>
<summary>★★ Expand ★★</summary>
```yaml
# Noise scale. May not do anything currently.
s_noise: 1.0
# ETA (basically ancestralness). May not do anything currently.
eta: 1.0
# Reversible ETA (used for reversible samplers). May not do anything currently.
reta: 1.0
# Sets the type of preview used for sampling in this group. One of:
# denoised: The default, shows the model prediction (takes positive and negative prompt into account).
# cond: Shows the model cond prediction (basically the positive prompt).
# uncond: Shows the model uncond prediction (basically the negative prompt).
# raw: Shows the raw noisy latent input.
preview_mode: denoised
# 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.
***
### `OCSNoise to SONAR_CUSTOM_NOISE`
Adapter that enables using OCS noise generators with nodes that accept `SONAR_CUSTOM_NOISE`.
Most built-in OCS nodes will accept either type currently.
***
### `OCSNoise PerlinSimple`
Generates 2D or 3D Perlin noise with many tuneable parameters. Can be plugged in to samplers for ancestral or SDE sampling. For initial noise or img2img workflows, use the `NoisyLatentLike` node from `ComfyUI-sonar` (see [Integration](#integration)).
3D Perlin noise works by taking a slice in the depth dimension each time the noise sampler is called.
For more tuneable parameters, see the `OCSNoise PerlinAdvanced` node.
**Note**: The shape of the latent must be a multiple of `lacunarity ** (octaves - 1) * res` (`**` indicates raising something to a power). Most latent types will have one latent pixel equaling eight normal pixels - i.e. if your image is 512x512, the latent would be 64x64.
#### Node Parameters
* `depth`: When non-zero, 3D perlin noise will be generated.
* `detail_level`: Controls the detail level of the noise when `break_pattern` is non-zero. No effect when using 100% raw Perlin noise.
* `octaves`: Generally controls the detail level of the noise. Each octave involves generating a layer of noise so there is a performance cost to increasing octaves.
* `persistence`: Controls how rough the generated noise is. Lower values will result in smoother noise, higher values will look more like Gaussian noise. Comma-separated list, multiple items will apply to octaves in sequence.
* `lacunarity`: Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.
* `res_height`: Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.
* `break_pattern`: Applies a function to break the Perlin pattern, making it more like normal noise. The value is the blend strength, where 1.0 indicates 100% pattern broken noise and 0.5 indicates 50% raw noise and 50% pattern broken noise. Generally should be at least 0.9 unless you want to generate colorful blobs.
***
### `OCSNoise PerlinAdvanced`
Generates 2D or 3D Perlin noise with many tuneable parameters. Can be plugged in to samplers for ancestral or SDE sampling. For initial noise or img2img workflows, use the `NoisyLatentLike` node from `ComfyUI-sonar` (see [Integration](#integration)).
3D Perlin noise works by taking a slice in the depth dimension each time the noise sampler is called.
**Note**: The shape of the latent in the relevant dimension _including padding_ must be a multiple of `lacunarity ** (octaves - 1) * res`. Most latent types will have one latent pixel equaling eight normal pixels - i.e. if your image is 512x512, the latent would be 64x64.
#### Node Parameters
* `depth`: When non-zero, 3D perlin noise will be generated.
* `detail_level`: Controls the detail level of the noise when `break_pattern` is non-zero. No effect when using 100% raw Perlin noise.
* `octaves`: Generally controls the detail level of the noise. Each octave involves generating a layer of noise so there is a performance cost to increasing octaves.
* `persistence`: Controls how rough the generated noise is. Lower values will result in smoother noise, higher values will look more like Gaussian noise. Comma-separated list, multiple items will apply to octaves in sequence.
* `lacunarity_height`: Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.
* `lacunarity_width`: " "
* `lacunarity_depth`: " "
* `res_height`: Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.
* `res_width`: " "
* `res_depth`: " "
* `break_pattern`: Applies a function to break the Perlin pattern, making it more like normal noise. The value is the blend strength, where 1.0 indicates 100% pattern broken noise and 0.5 indicates 50% raw noise and 50% pattern broken noise. Generally should be at least 0.9 unless you want to generate colorful blobs.
* `initial_depth`: First zero-based depth index the noise generator will return. Only has an effect when depth is non-zero.
* `wrap_depth`: If non-zero, instead of generating a new chunk of noise when the last slice is used will instead jump back to the specified zero-based depth index. Only has an effect when depth is non-zero. Since this is repeating the same noise, you may need to reduce `s_noise` in samplers especially if your `depth` value is low.
* `max_depth`: Basically crops the depth dimension to the specified value (inclusive). Negative values start from the end, the default of -1 does no cropping. Only has an effect when depth is non-zero. The reason you might want to use this is changing `depth` will also effectively change the seed.
* `tileable_height`: Makes the specified dimension tileable. (May or may not work correctly.)
* `tileable_width`: " "
* `tileable_depth`: " "
* `blend`: Blending function used when generating Perlin noise. When set to values other than LERP may not work at all or may not actually generate Perlin noise. If you have `ComfyUI-bleh` there will be many more blending options (see [Integration](#integration)).
* `pattern_break_blend`: Blending function used to blend pattern broken noise with raw noise. If you have `ComfyUI-bleh` there will be many more blending options (see [Integration](#integration)).
* `depth_over_channels`: When disabled, each channel will have its own separate 3D noise pattern. When enabled, depth is multiplied by the number of channels and each channel is a slice of depth. Only has an effect when depth is non-zero.
* `pad_height`: Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation.
* `pad_width`: " "
* `pad_depth`: " "
* `initial_amplitude`: Controls the amplitude for the first octave. The amplitude gets multiplied by `persistence` after each octave.
* `initial_frequency_height`: Controls the frequency for the first octave for the this axis. The frequency gets multiplied by `lacunarity` after each octave.
* `initial_frequency_width`: " "
* `initial_frequency_depth`: " "
* `normalize`: Controls whether the output noise is normalized after generation.
* `device`: Controls what device is used to generate the noise. GPU noise may be slightly faster but you will get different results on different GPUs.
+9 -4
View File
@@ -1,8 +1,13 @@
from .py import nodes
from .py import custom_noise
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,
} | custom_noise.NODE_CLASS_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS"]
Binary file not shown.

After

Width:  |  Height:  |  Size: 48 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 31 KiB

+164
View File
@@ -0,0 +1,164 @@
# 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.
`:=` is used to assign to a temporary variable (see `set_var` below).
The expression language supports a C/JavaScript style ternary operator: `condition ? true_branch : false_branch` is the equivalent of `if(condition, true_branch, false_branch)`.
Like Python, a parenthesized expression with a trailing comma can be used to create an empty tuple. Example: `(1,)`
## Filter Variables
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> |
|⬤| `set_var` | `SY`, `*` | `*` |
| <td colspan=3 align=left>Sets a temporary variable to the specified value and returns the value. Alias for the `:=` assignment operator. <br/> **Example**: `test1 := 2; set_var('test2, 10); test1 * test2`</td> |
|⬤| `unsafe_call` | `callable`, `*`\* | `*` |
| <td colspan=3 align=left>Allows calling an arbitrary callable. <br/> **Example:** `unsafe_call(some_callable, arg1, arg2, kwarg1 :> 123)`</td>
**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_scale_nnlatentupscale` | tensor:`T`, mode:`SY(sd1)`, scale:`SN(2.0)` | `T` |
| <td colspan=3 align=left>Available if you have [ComfyUi_NNLatentUpscale](https://github.com/Ttl/ComfyUi_NNLatentUpscale) installed. `mode` must be one of `sd1` or `sdxl`. `scale` should be between 1.0 and 2.0 (may or may not work out of that range).<br/> **Example:** `t_scale_nnlatentupscale(some_tensor, 'sdxl, 1.5)`</td> |
|⬤| `t_shape` | tensor:`T` | `SN` |
| <td colspan=3 align=left>Returns a tensor's shape as a tuple. <br/> Example: `shp := t_shape(some_tensor); width := shp[-1]; height := shp[-2]`</td> |
|⬤| `t_std` | tensor:`T`, dim:`SN(-3, -2, -1)` | `T` |
| <td colspan=3 align=left>Tensor std, second argument is dimensions. <br/> **Example:** `t_std(some_tensor, (-2, -1))`</td> |
|⬤| `t_taesd_decode` | tensor:`T`, mode:`SY(sd15)` | `T` |
| <td colspan=3 align=left>Decodes a latent tensor used TAESD. Mode must be one of `sd15`, `sdxl`. Only works if the appropriate models are in `vae_approx` <br/> **Example:** `t_taesd_decode(some_tensor, 'sd15)`</td> |
|⬤| `unsafe_tensor_method` | `T`, `SY`, `*`\* | `*` |
| <td colspan=3 align=left>Unsafe tensor method call. See note below. <br/> **Example:** `unsafe_tensor_method(some_tensor, 'mul, 10)`</td> |
|⬤| `unsafe_torch` | path:`SY` | `*` |
| <td colspan=3 align=left>Unsafe Torch module attribute access. See note below. <br/> **Example:** `unsafe_torch('nn.functional.interpolate)`</td> |
**Note on `unsafe_tensor_method` and `unsafe_torch`**: These functions are disabled by default. If the environment variable `COMFYUI_OCS_ALLOW_UNSAFE_EXPRESSIONS` is set to anything then you can use `unsafe_tensor_method` with a whitelisted set of methods (best effort to avoid anything actually unsafe). If the environment variable `COMFYUI_OCS_ALLOW_ALL_UNSAFE` is set to anything then `unsafe_torch` is enabled and `unsafe_tensor_method` will allow calling any method. ***WARNING***: Allowing _all_ unsafe with workflows you don't trust is _not_ recommended and a malicious workflow will likely have access to anything ComfyUI can access. It is effectively the same as letting the workflow run an arbitrary script on your system.
## Tensor Expression Functions
`IMG` used here to donate the type for functions that take an image. This may actually be an image batch rather than a single image.
| | Name | Input | Output |
| :--- | :--- | :--- | :--- |
|⬤| `img_pil_resize` | image:`IMG`, size:`SN \| NS`, resample_mode:`SY(bicubic)`, absolute_scale:`B(false)` | `IMG` |
| <td colspan=3 align=left>Scales an image batch using Pillow's [`Image.resize`](https://pillow.readthedocs.io/en/stable/reference/Image.html#PIL.Image.Image.resize) function (follow link for information about resample modes, etc). If scale is a tuple, it will be interpreted as `(height, width)`. When `absolute_scale` is not set, the scales will be interpreted as percentages otherwise absolute values will be used. <br/> Example: `img_pil_resize(image_batch, (0.75, 0.5), 'lanczos)`</td> |
|⬤| `img_shape` | image:`IMG` | `SN` |
| <td colspan=3 align=left>Returns an image's shape as a tuple. Will fail if all the images in the batch aren't the same size. <br/> Example: `shp := img_shape(image_batch); width := shp[-1]; height := shp[-2]`</td> |
|⬤| `img_taesd_encode` | image:`IMG`, reference_latent: `T`, mode:`SY(sd15)` | `T` |
| <td colspan=3 align=left>Encodes an image batch into a latent tensor. The reference latent is only used to determine what device and type the output should be. Mode must be one of `sd15`, `sdxl`. Only works if the appropriate models are in `vae_approx`. <br/> **Example:** `img_taesd_encode(image_batch, some_tensor, 'sd15)`</td> |
+189
View File
@@ -0,0 +1,189 @@
# 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.
If you have [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) available, you can use any blend mode it supports. Otherwise OCS provides these built-in blend modes: `lerp`, `a_only`, `b_only`. _Note_: `a` is considered the original value, `b` the changed value. `a_only` and `b_only` will still scale their output by the `strength`.
## Filter Types
### `simple`
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
```
+8
View File
@@ -0,0 +1,8 @@
from . import noise_perlin
from . import nodes
NODE_CLASS_MAPPINGS = {
"OCSNoise PerlinSimple": noise_perlin.PerlinSimpleNode,
"OCSNoise PerlinAdvanced": noise_perlin.PerlinAdvancedNode,
"OCSNoise to SONAR_CUSTOM_NOISE": nodes.ToSonarNode,
}
+178
View File
@@ -0,0 +1,178 @@
import abc
import torch
from typing import Callable, Any
from ..noise import scale_noise
class CustomNoiseItemBase(abc.ABC):
def __init__(self, factor, **kwargs):
self.factor = factor
self.keys = set(kwargs.keys())
for k, v in kwargs.items():
setattr(self, k, v)
def clone_key(self, k):
return getattr(self, k)
def clone(self):
return self.__class__(self.factor, **{k: self.clone_key(k) for k in self.keys})
def set_factor(self, factor):
self.factor = factor
return self
def get_normalize(self, k, default=None):
val = getattr(self, k, None)
return default if val is None else val
@abc.abstractmethod
def make_noise_sampler(
self,
x: torch.Tensor,
sigma_min=None,
sigma_max=None,
seed=None,
cpu=True,
normalized=True,
):
raise NotImplementedError
class CustomNoiseChain:
def __init__(self, items=None):
self.items = items if items is not None else []
def clone(self):
return CustomNoiseChain(
[i.clone() for i in self.items],
)
def add(self, item):
if item is None:
raise ValueError("Attempt to add nil item")
self.items.append(item)
@property
def factor(self):
return sum(abs(i.factor) for i in self.items)
def rescaled(self, scale=1.0):
divisor = self.factor / scale
divisor = divisor if divisor != 0 else 1.0
result = self.clone()
if divisor != 1:
for i in result.items:
i.set_factor(i.factor / divisor)
return result
@torch.no_grad()
def make_noise_sampler(
self,
x: torch.Tensor,
sigma_min=None,
sigma_max=None,
seed=None,
cpu=True,
normalized=True,
) -> Callable:
noise_samplers = tuple(
i.make_noise_sampler(
x,
sigma_min,
sigma_max,
seed=seed,
cpu=cpu,
normalized=False,
)
for i in self.items
)
if not noise_samplers or not all(noise_samplers):
raise ValueError("Failed to get noise sampler")
factor = self.factor
def noise_sampler(sigma, sigma_next):
result = None
for ns in noise_samplers:
noise = ns(sigma, sigma_next)
if result is None:
result = noise
else:
result += noise
return scale_noise(result, factor, normalized=normalized)
return noise_sampler
class CustomNoiseNodeBase(abc.ABC):
DESCRIPTION = "An Overly Complicated Sampling custom noise item."
RETURN_TYPES = ("OCS_NOISE",)
OUTPUT_TOOLTIPS = ("A custom noise chain.",)
CATEGORY = "OveryComplicatedSampling/noise"
FUNCTION = "go"
@abc.abstractmethod
def get_item_class(self):
raise NotImplementedError
@classmethod
def INPUT_TYPES(cls, *, include_rescale=True, include_chain=True):
result = {
"required": {
"factor": (
"FLOAT",
{
"default": 1.0,
"min": -100.0,
"max": 100.0,
"step": 0.001,
"round": False,
"tooltip": "Scaling factor for the generated noise of this type.",
},
),
},
"optional": {},
}
if include_rescale:
result["required"] |= {
"rescale": (
"FLOAT",
{
"default": 0.0,
"min": 0.0,
"max": 100.0,
"step": 0.001,
"round": False,
"tooltip": "When non-zero, this custom noise item and other custom noise items items connected to it will have their factor scaled to add up to the specified rescale value.",
},
),
}
if include_chain:
result["optional"] |= {
"ocs_noise_opt": (
"OCS_NOISE",
{
"tooltip": "Optional input for more custom noise items.",
},
),
}
return result
def go(
self,
factor=1.0,
rescale=0.0,
ocs_noise_opt=None,
**kwargs: dict[str, Any],
):
nis = ocs_noise_opt.clone() if ocs_noise_opt else CustomNoiseChain()
if factor != 0:
nis.add(self.get_item_class()(factor, **kwargs))
return (nis if rescale == 0 else nis.rescaled(rescale),)
class NormalizeNoiseNodeMixin:
@staticmethod
def get_normalize(val: str) -> None | bool:
return None if val == "default" else val == "forced"
+16
View File
@@ -0,0 +1,16 @@
class ToSonarNode:
RETURN_TYPES = ("SONAR_CUSTOM_NOISE",)
CATEGORY = "OveryComplicatedSampling/noise"
FUNCTION = "go"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ocs_noise": ("OCS_NOISE",),
},
}
@classmethod
def go(cls, ocs_noise):
return (ocs_noise,)
+744
View File
@@ -0,0 +1,744 @@
import torch
import math
import itertools
from .base import CustomNoiseItemBase, CustomNoiseNodeBase, NormalizeNoiseNodeMixin
from ..latent import normalize_to_scale
from ..noise import scale_noise
from ..filtering import BLENDING_MODES
# Perlin generation routines based on https://github.com/Extraltodeus/noise_latent_perlinpinpin which was based on https://gist.github.com/vadimkantorov/ac1b097753f217c5c11bc2ff396e0a57 which was based on https://github.com/pvigier/perlin-numpy
def smoothstep_function(t):
return 6 * t**5 - 15 * t**4 + 10 * t**3
class DEFAULTS:
depth = 16
res = ((1,), (1,), (1,))
octaves = 2
persistence = (1.0,)
lacunarity = ((2,), (2,), (2,))
initial_amplitude = 1.0
initial_frequency = (1.0, 1.0, 1.0)
break_pattern = 1.0
detail_level = 0.0
tileable = (False, False, False)
fade = smoothstep_function
blend = "lerp"
pattern_break_blend = "lerp"
depth_over_channels = False
initial_depth = 0
wrap_depth = 0
max_depth = -1
pad = (0, 0, 0)
generator = None
device = "default"
@classmethod
def get_commasep(cls, key, idx=None):
val = getattr(cls, key)
if idx is not None:
val = val[idx]
return ", ".join(repr(v) for v in val)
def rand_perlin(
shape,
res,
*,
tileable=DEFAULTS.tileable,
fade=DEFAULTS.fade,
blend=BLENDING_MODES[DEFAULTS.blend],
generator=DEFAULTS.generator,
device=DEFAULTS.device,
):
dims = len(res)
didxs = tuple(range(dims))
delta, d = zip(*((res[i] / shape[i], int(shape[i] // res[i])) for i in didxs))
grid = (
torch.stack(
torch.meshgrid(*(torch.arange(0, res[i], delta[i]) for i in didxs)),
dim=-1,
)
% 1
).to(device=device)
noise = (
2
* math.pi
* torch.rand(
max(1, dims - 1),
*(round(res[i]) + 1 for i in didxs),
generator=generator,
device=device,
)
)
if dims == 1:
gradients = torch.cos(noise[0])
elif dims == 2:
gradients = torch.stack((torch.cos(noise[0]), torch.sin(noise[0])), dim=-1)
elif dims == 3:
gradients = torch.stack(
(
torch.sin(noise[0]) * torch.cos(noise[1]),
torch.sin(noise[0]) * torch.sin(noise[1]),
torch.cos(noise[0]),
),
dim=-1,
)
elif dims == 4:
# No idea if this makes sense.
gradients = torch.stack(
(
torch.sin(noise[0]) * torch.cos(noise[1]),
torch.sin(noise[0]) * torch.sin(noise[1]),
torch.sin(noise[1]) * torch.cos(noise[2]),
torch.sin(noise[1]) * torch.sin(noise[2]),
),
dim=-1,
)
else:
raise ValueError("Currently only dimensions up to 4 are supported")
del noise
for tidx, tile in enumerate(tileable[:dims]):
if not tile:
continue
gradients[tuple(-1 if didx == tidx else None for didx in didxs)] = gradients[
tuple(0 if didx == tidx else None for didx in didxs)
]
shape_slices = tuple(slice(0, shape[i]) for i in didxs)
def tile_grads(slices):
result = gradients[*(slice(*slices[i]) for i in didxs)]
for i in didxs:
result = result.repeat_interleave(d[i], i)
return result
def dot(grad, shift):
return (
torch.stack(tuple(grid[*shape_slices, i] + shift[i] for i in didxs), dim=-1)
* grad[shape_slices]
).sum(dim=-1)
# It's just binary with the bits reversed and -1 for enabled columns.
def get_shift(n, dims, *, on_value, off_value):
return tuple(
on_value if n & (1 << bitidx) else off_value for bitidx in range(dims)
)
def blend_reduce(vals, t, depth=0):
curr_t = t[..., depth]
if len(vals) == 2:
return blend(*vals, curr_t)
return blend_reduce(
tuple(blend(v1, v2, curr_t) for v1, v2 in itertools.batched(vals, 2)),
t,
depth + 1,
)
ns = tuple(
dot(
tile_grads(get_shift(i, dims, off_value=(None, -1), on_value=(1, None))),
get_shift(i, dims, off_value=0, on_value=-1),
)
for i in range(1 << dims)
)
return math.sqrt(2) * blend_reduce(ns, fade(grid[shape_slices]))
def generate_fractal_noise(
shape,
res=DEFAULTS.res,
octaves=DEFAULTS.octaves,
persistence=DEFAULTS.persistence,
lacunarity=DEFAULTS.lacunarity,
initial_amplitude=DEFAULTS.initial_amplitude,
initial_frequency=DEFAULTS.initial_frequency,
tileable=DEFAULTS.tileable,
fade=DEFAULTS.fade,
blend=BLENDING_MODES[DEFAULTS.blend],
generator=DEFAULTS.generator,
device=DEFAULTS.device,
):
ndim = len(shape)
def get_wrap_dim(val, *dims):
for dim in dims:
nelem = len(val) if not isinstance(val, torch.Tensor) else val.shape[0]
val = val[dim % nelem]
return val
def get_unwrapped_octaves_dims(val):
return torch.tensor(
tuple(
get_wrap_dim(val, didx, oidx)
for oidx in range(octaves)
for didx in range(ndim)
),
dtype=torch.float,
device="cpu",
).view(octaves, ndim)
res = get_unwrapped_octaves_dims(res)
lacunarity = get_unwrapped_octaves_dims(lacunarity)
initial_frequency = initial_frequency[-ndim:]
persistence = persistence[:octaves]
noise = torch.zeros(shape, dtype=torch.float32, device=device)
frequency = torch.ones(ndim, dtype=torch.float, device="cpu")
frequency[: len(initial_frequency)] = frequency.new(initial_frequency)
amplitude = initial_amplitude
for octave in range(octaves):
noise += amplitude * rand_perlin(
shape,
tuple(
frequency[didx].item() * res[octave][didx].item()
for didx in range(ndim)
),
tileable=tileable,
fade=fade,
blend=blend,
generator=generator,
device=device,
)
# print(
# f"Octave {octave}: freq={frequency}, amp={amplitude}, lac={lacunarity[octave]}, pers={get_wrap_dim(persistence, octave)}"
# )
frequency *= lacunarity[octave]
amplitude *= get_wrap_dim(persistence, octave)
# print(f"Octave {octave}: POST: freq={frequency}, amp={amplitude}")
return noise
def create_noisy_latents_perlin(
width,
height,
depth,
*,
batch_size=1,
detail_level=DEFAULTS.detail_level,
octaves=DEFAULTS.octaves,
persistence=DEFAULTS.persistence,
lacunarity=DEFAULTS.lacunarity,
tileable=DEFAULTS.tileable,
res=DEFAULTS.res,
break_pattern=DEFAULTS.break_pattern,
channels=4,
blend=BLENDING_MODES[DEFAULTS.blend],
pattern_break_blend=BLENDING_MODES[DEFAULTS.pattern_break_blend],
depth_over_channels=DEFAULTS.depth_over_channels,
pad=DEFAULTS.pad,
initial_frequency=DEFAULTS.initial_frequency,
initial_amplitude=DEFAULTS.initial_amplitude,
generator=DEFAULTS.generator,
device=DEFAULTS.device,
):
pad_depth, pad_height, pad_width = pad
if depth < 1:
depth_over_channels = False
pad_depth = 0
shape = (height, width)
eff_shape = (
height + pad_height * 2,
width + pad_width * 2,
)
eff_channels = channels if not depth_over_channels else 1
eff_depth = depth if not depth_over_channels else depth * channels
if depth > 0:
shape = (depth, height, width)
eff_shape = (
eff_depth + pad_depth * 2,
height + pad_height * 2,
width + pad_width * 2,
)
noise = torch.zeros(
(batch_size, channels, *shape),
dtype=torch.float32,
device=device,
)
noise_dims = len(shape)
for i in range(batch_size):
for j in range(eff_channels):
noise_values = generate_fractal_noise(
eff_shape,
res=res,
octaves=octaves,
persistence=persistence,
lacunarity=lacunarity,
tileable=tileable,
blend=blend,
initial_frequency=initial_frequency,
initial_amplitude=initial_amplitude,
generator=generator,
device=device,
)
noise_values = normalize_to_scale(noise_values, -1.0, 1.0, dim=())
if break_pattern != 0:
result = torch.remainder(torch.abs(noise_values) * 1000000, 11) / 11
result = (
((1 + detail_level / 10) * torch.erfinv(2 * result - 1) * (2**0.5))
.mul_(0.2)
.clamp_(-1, 1)
)
result = pattern_break_blend(noise_values, result, break_pattern)
else:
result = noise_values
if pad_width + pad_height + pad_depth > 0:
result = (
result[
...,
pad_depth : eff_depth + pad_depth,
pad_height : height + pad_height,
pad_width : width + pad_width,
]
if noise_dims == 3
else result[
...,
pad_height : height + pad_height,
pad_width : width + pad_width,
]
)
if not depth_over_channels:
noise[i, j, ...] = result
continue
noise[i, ...] = result.view(depth, channels, height, width).movedim(0, 1)
return noise.movedim(-3, 0) if noise_dims == 3 else noise
class PerlinItem(CustomNoiseItemBase):
def __init__(
self,
factor,
*,
depth=20,
detail_level=DEFAULTS.detail_level,
octaves=DEFAULTS.octaves,
persistence=DEFAULTS.persistence,
lacunarity_depth=DEFAULTS.lacunarity[0],
lacunarity_height=DEFAULTS.lacunarity[1],
lacunarity_width=DEFAULTS.lacunarity[2],
lacunarity=None,
tileable_depth=DEFAULTS.tileable[0],
tileable_height=DEFAULTS.tileable[1],
tileable_width=DEFAULTS.tileable[2],
tileable=None,
res_depth=DEFAULTS.res[0],
res_height=DEFAULTS.res[1],
res_width=DEFAULTS.res[2],
res=None,
initial_frequency_depth=DEFAULTS.initial_frequency[0],
initial_frequency_height=DEFAULTS.initial_frequency[1],
initial_frequency_width=DEFAULTS.initial_frequency[2],
initial_frequency=None,
initial_amplitude=DEFAULTS.initial_amplitude,
wrap_depth=DEFAULTS.wrap_depth,
initial_depth=DEFAULTS.initial_depth,
max_depth=DEFAULTS.max_depth,
break_pattern=DEFAULTS.break_pattern,
blend=DEFAULTS.blend,
pattern_break_blend=DEFAULTS.pattern_break_blend,
depth_over_channels=DEFAULTS.depth_over_channels,
pad=None,
pad_depth=DEFAULTS.pad[0],
pad_height=DEFAULTS.pad[1],
pad_width=DEFAULTS.pad[2],
device=None,
normalized=None,
**kwargs,
):
if tileable is None:
tileable = (tileable_depth, tileable_height, tileable_width)[
int(depth == 0) :
]
if res is None:
res = self.maybe_parse_dhw_triple(
(res_depth, res_height, res_width), depth, int
)
if lacunarity is None:
lacunarity = self.maybe_parse_dhw_triple(
(
lacunarity_depth,
lacunarity_height,
lacunarity_width,
),
depth,
)
if pad is None:
pad = (pad_depth, pad_height, pad_width)
if initial_frequency is None:
initial_frequency = (
initial_frequency_depth,
initial_frequency_height,
initial_frequency_width,
)[int(depth == 0) :]
persistence = self.maybe_parse_commasep_list(persistence)
super().__init__(
factor,
depth=depth,
detail_level=detail_level,
octaves=octaves,
persistence=persistence,
lacunarity=lacunarity,
tileable=tileable,
res=res,
initial_frequency=initial_frequency,
initial_amplitude=initial_amplitude,
initial_depth=initial_depth,
wrap_depth=wrap_depth,
max_depth=max_depth,
break_pattern=break_pattern,
blend=blend,
pattern_break_blend=pattern_break_blend,
depth_over_channels=depth_over_channels,
pad=pad,
device=device,
normalized=normalized
if not isinstance(normalized, str)
else NormalizeNoiseNodeMixin.get_normalize(normalized),
**kwargs,
)
@classmethod
def maybe_parse_dhw_triple(cls, val, depth, convert=float):
return tuple(cls.maybe_parse_commasep_list(v) for v in val)[int(depth == 0) :]
@classmethod
def maybe_parse_commasep_list(cls, val, convert=float):
if not isinstance(val, str):
return val
return tuple(convert(v) for v in val.strip().split(",") if v.strip())
def make_noise_sampler(
self,
x: torch.Tensor,
sigma_min: float | None,
sigma_max: float | None,
seed: int | None,
cpu: bool = True,
normalized=True,
):
normalized = self.get_normalize("normalized", normalized)
cpu = cpu if self.device == "default" else self.device == "cpu"
device = torch.device(0 if not cpu else "cpu")
noise_chunk = None
noise_index = self.initial_depth
max_idx = None
b, c, h, w = x.shape
x_device, x_dtype = x.device, x.dtype
del x
blend = BLENDING_MODES[self.blend]
pattern_break_blend = BLENDING_MODES[self.pattern_break_blend]
def noise_sampler(_s, _sn):
nonlocal noise_chunk, noise_index, max_idx
if noise_chunk is None:
# print("-->", noise_index, self.depth)
noise_chunk = create_noisy_latents_perlin(
w,
h,
self.depth,
batch_size=b,
channels=c,
detail_level=self.detail_level,
octaves=self.octaves,
persistence=self.persistence,
lacunarity=self.lacunarity,
initial_frequency=self.initial_frequency,
initial_amplitude=self.initial_amplitude,
break_pattern=self.break_pattern,
res=self.res,
tileable=self.tileable,
blend=blend,
pattern_break_blend=pattern_break_blend,
depth_over_channels=self.depth_over_channels,
pad=self.pad,
device=device,
).to(device=x_device, dtype=x_dtype)
if self.depth < 1: # 2D mode
noise = noise_chunk
noise_chunk = None
return scale_noise(noise, self.factor, normalized=normalized)
if self.max_depth != 0 and self.max_depth != -1:
noise_chunk = noise_chunk[: self.max_depth]
chunk_shape = noise_chunk.shape
max_idx = (
chunk_shape[0] - 1
if self.wrap_depth == 0
else min(self.wrap_depth, chunk_shape[0] - 1)
)
if max_idx < 0:
max_idx += chunk_shape[0]
noise = noise_chunk[noise_index]
noise_index += 1
if noise_index > max_idx:
noise_index = 0
if not self.wrap_depth:
noise_chunk = None
return scale_noise(noise, self.factor, normalized=normalized)
return noise_sampler
class PerlinAdvancedNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin):
DESCRIPTION = "Advanced Perlin noise generator, allows generating 2D or 3D Perlin noise. See the OCSNoise PerlinSimple node for less tuneable parameters."
@classmethod
def INPUT_TYPES(cls):
result = super().INPUT_TYPES()
result["required"] |= {
"depth": (
"INT",
{
"default": DEFAULTS.depth,
"tooltip": "When non-zero, 3D perlin noise will be generated.",
},
),
"detail_level": (
"FLOAT",
{
"default": DEFAULTS.detail_level,
"tooltip": "Controls the detail level of the noise when break_pattern is non-zero. No effect when using 100% raw Perlin noise.",
},
),
"octaves": (
"INT",
{
"default": DEFAULTS.octaves,
"tooltip": "Generally controls the detail level of the noise. Each octave involves generating a layer of noise so there is a performance cost to increasing octaves.",
},
),
"persistence": (
"STRING",
{
"default": DEFAULTS.get_commasep("persistence"),
"tooltip": "Controls how rough the generated noise is. Lower values will result in smoother noise, higher values will look more like Gaussian noise. Comma-separated list, multiple items will apply to octaves in sequence.",
},
),
"lacunarity_height": (
"STRING",
{
"default": DEFAULTS.get_commasep("lacunarity", 0),
"tooltip": "Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.",
},
),
"lacunarity_width": (
"STRING",
{
"default": DEFAULTS.get_commasep("lacunarity", 1),
"tooltip": "Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.",
},
),
"lacunarity_depth": (
"STRING",
{
"default": DEFAULTS.get_commasep("lacunarity", 2),
"tooltip": "Lacunarity controls the frequency multiplier between successive octaves. Only has an effect when depth is non-zero and octaves is greater than one. Comma-separated list, multiple items will apply to octaves in sequence.",
},
),
"res_height": (
"STRING",
{
"default": DEFAULTS.get_commasep("res", 0),
"tooltip": "Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.",
},
),
"res_width": (
"STRING",
{
"default": DEFAULTS.get_commasep("res", 1),
"tooltip": "Number of periods of noise to generate along an axis. Comma-separated list, multiple items will apply to octaves in sequence.",
},
),
"res_depth": (
"STRING",
{
"default": DEFAULTS.get_commasep("res", 2),
"tooltip": "Number of periods of noise to generate along an axis. Only has an effect when depth is non-zero. Comma-separated list, multiple items will apply to octaves in sequence.",
},
),
"break_pattern": (
"FLOAT",
{
"default": DEFAULTS.break_pattern,
"tooltip": "Applies a function to break the Perlin pattern, making it more like normal noise. The value is the blend strength, where 1.0 indicates 100% pattern broken noise and 0.5 indicates 50% raw noise and 50% pattern broken noise. Generally should be at least 0.9 unless you want to generate colorful blobs.",
},
),
"initial_depth": (
"INT",
{
"default": DEFAULTS.initial_depth,
"tooltip": "First zero-based depth index the noise generator will return. Only has an effect when depth is non-zero.",
},
),
"wrap_depth": (
"INT",
{
"default": DEFAULTS.wrap_depth,
"tooltip": "If non-zero, instead of generating a new chunk of noise when the last slice is used will instead jump back to the specified zero-based depth index. Only has an effect when depth is non-zero.",
},
),
"max_depth": (
"INT",
{
"default": DEFAULTS.max_depth,
"tooltip": "Basically crops the depth dimension to the specified value (inclusive). Negative values start from the end, the default of -1 does no cropping. Only has an effect when depth is non-zero.",
},
),
"tileable_height": (
"BOOLEAN",
{
"default": DEFAULTS.tileable[0],
"tooltip": "Makes the specified dimension tileable.",
},
),
"tileable_width": (
"BOOLEAN",
{
"default": DEFAULTS.tileable[1],
"tooltip": "Makes the specified dimension tileable.",
},
),
"tileable_depth": (
"BOOLEAN",
{
"default": DEFAULTS.tileable[2],
"tooltip": "Makes the specified dimension tileable. Only has an effect when depth is non-zero.",
},
),
"blend": (
tuple(BLENDING_MODES.keys()),
{
"default": "lerp",
"tooltip": "Blending function used when generating Perlin noise. When set to values other than LERP may not work at all or may not actually generate Perlin noise.",
},
),
"pattern_break_blend": (
tuple(BLENDING_MODES.keys()),
{
"default": "lerp",
"tooltip": "Blending function used to blend pattern broken noise with raw noise.",
},
),
"depth_over_channels": (
"BOOLEAN",
{
"default": DEFAULTS.depth_over_channels,
"tooltip": "When disabled, each channel will have its own separate 3D noise pattern. When enabled, depth is multiplied by the number of channels and each channel is a slice of depth. Only has an effect when depth is non-zero.",
},
),
"pad_height": (
"INT",
{
"default": DEFAULTS.pad[0],
"min": 0,
"tooltip": "Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation.",
},
),
"pad_width": (
"INT",
{
"default": DEFAULTS.pad[1],
"min": 0,
"tooltip": "Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation.",
},
),
"pad_depth": (
"INT",
{
"default": DEFAULTS.pad[2],
"min": 0,
"tooltip": "Pads the specified dimension by the size. Equal padding will be added on both sides and cropped out after generation. Only has an effect when depth is non-zero.",
},
),
"initial_amplitude": (
"FLOAT",
{
"default": DEFAULTS.initial_amplitude,
"tooltip": "Controls the amplitude for the first octave.",
},
),
"initial_frequency_height": (
"FLOAT",
{
"default": DEFAULTS.initial_frequency[0],
"tooltip": "Controls the frequency for the first octave for the this axis.",
},
),
"initial_frequency_width": (
"FLOAT",
{
"default": DEFAULTS.initial_frequency[1],
"tooltip": "Controls the frequency for the first octave for the this axis.",
},
),
"initial_frequency_depth": (
"FLOAT",
{
"default": DEFAULTS.initial_frequency[2],
"tooltip": "Controls the frequency for the first octave for the this axis.",
},
),
"normalize": (
("default", "forced", "off"),
{
"tooltip": "Controls whether the output noise is normalized after generation.",
},
),
"device": (
("default", "cpu", "gpu"),
{
"default": "default",
"tooltip": "Controls what device is used to generate the noise. GPU noise may be slightly faster but you will get different results on different GPUs.",
},
),
}
return result
@classmethod
def get_item_class(cls):
return PerlinItem
class PerlinSimpleNode(PerlinAdvancedNode):
DESCRIPTION = "Simplified Perlin noise generator, allows generating 2D or 3D Perlin noise. See the OCSNoise PerlinAdvanced node for more tuneable parameters."
_COPY_KEYS = {
"factor",
"rescale",
"depth",
"detail_level",
"octaves",
"persistence",
"break_pattern",
}
@classmethod
def INPUT_TYPES(cls):
result = super().INPUT_TYPES()
orig_reqs = result["required"]
reqs = {k: v for k, v in orig_reqs.items() if k in cls._COPY_KEYS}
reqs["lacunarity"] = orig_reqs["lacunarity_height"]
reqs["res"] = orig_reqs["res_height"]
result["required"] = reqs
return result
@classmethod
def get_item_class(cls):
def wrapper(factor, *, lacunarity, res, **kwargs):
return PerlinItem(
factor,
lacunarity_height=lacunarity,
lacunarity_width=lacunarity,
lacunarity_depth=lacunarity,
res_height=res,
res_width=res,
res_depth=res,
**kwargs,
)
return wrapper
+19
View File
@@ -0,0 +1,19 @@
from . import types, expression, handler, util, validation
from .expression import Expression
from .validation import Arg, ValidateArg
from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext
__all__ = (
"Arg",
"BaseHandler",
"BASIC_HANDLERS",
"expression",
"Expression",
"handler",
"HandlerContext",
"types",
"util",
"ValidateArg",
"validation",
)
+265
View File
@@ -0,0 +1,265 @@
import re
import operator
from .parser import Parser, ParserSpec, ParseError
from .types import (
Empty,
ExpBase,
ExpOp,
ExpBinOp,
ExpSym,
ExpStatements,
ExpFunAp,
ExpTuple,
ExpDict,
ExpKV,
)
COMMA_PRECEDENCE = 2
class Expression:
EXPR_RE = re.compile(
r"""
\s*
(
\d+ # Numeric literal
(?: \. \d* )? # Floating point
(?: e [+-] \d+)? # Scientific notation
| (?: \*\* | // ) # Doubled operators
| [<>]=? # Relative comparison
| [!=]= # Equality
| (?: \|\| | && ) # Logic
| [-+*/|!(),] # Operators
| :> # Key value binop
| := # Assignment
| ; # Sequencing
| [?:] # Ternary
| \[ | ] # Index
| \.\.\. # Index ellipsis
| '[-\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):
if self.expr != ExpOp("default"):
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))
CONST_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,
"neg": operator.neg,
">": operator.gt,
"<": operator.lt,
">=": operator.ge,
"<=": operator.le,
"!=": operator.ne,
"==": operator.eq,
}
def is_const_value(val):
return val in (None, True, False) or isinstance(val, (int, float, ExpSym))
def make_funap(op, args=(), kwargs=None):
if kwargs is None:
kwargs = ExpDict()
argc = len(args)
if argc > 2 or len(kwargs) or not all(is_const_value(v) for v in args):
return ExpFunAp(op, args, kwargs)
if argc == 1 and op in "-+":
return -args[0] if op == "-" else args[0]
h = CONST_OP_HANDLERS.get(op)
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(COMMA_PRECEDENCE))
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):
r = None if p.token in (None, ")", ";") else p.parse_until(0)
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 left_assign(p, token, left, bp):
if not isinstance(left, (ExpOp, ExpSym)):
raise ParseError(f"bad LHS type for assignment operation {type(left)}")
val = p.parse_until(bp)
return make_funap("set_var", ExpTuple((ExpSym(left), val)))
@staticmethod
def left_ternary(p, token, left, bp):
true_branch = p.parse_until(0)
p.expect(":")
false_branch = p.parse_until(bp)
return make_funap("if", ExpTuple((left, true_branch, false_branch)))
@staticmethod
def get_type(token):
if isinstance(token, (int, float)):
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_leftright(5, self.left_ternary, ("?",))
self.add_leftright(4, self.left_assign, (":=",))
self.add_left(COMMA_PRECEDENCE, self.left_comma, (",",))
self.add_left(1, self.left_semicolon, (";",))
self.add_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, (")", "]", ":"))
+415
View File
@@ -0,0 +1,415 @@
import operator
from .validation import ValidateArg, Arg, ValidateError
from .types import Empty, ExpDict, ExpOp
from .util import torch
class HandlerError(Exception):
pass
class HandlerContext:
def __init__(self, handlers=None, constants=None, variables=None):
self.handlers = handlers if handlers is not None else {}
self.constants = constants if constants is not None else {}
self.variables = variables if variables is not None else {}
def get_handler(self, k, default=Empty):
return self.handlers.get(k, default)
def get_var(self, k, default=Empty):
result = self.constants.get(k, Empty)
if result is Empty:
result = self.variables.get(k, Empty)
return default if result is Empty else result
def set_var(self, k, v):
if k in self.constants:
raise KeyError(
f"Cannot set variable with key {k}: already exists as a constant"
)
self.variables[k] = v
def unset_var(self, k):
if k in self.variables:
del self.variables[k]
return True
return False
def __contains__(self, k):
return any(
k in coll for coll in (self.handlers, self.constants, self.variables)
)
def clone(self, *, handlers=Empty, constants=Empty, variables=Empty):
return self.__class__(
self.handlers if handlers is Empty else handlers,
self.constants if constants is Empty else constants,
self.variables if variables is Empty else variables,
)
class BaseHandler:
input_validators = ()
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} for {obj.name}, 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} for {obj.name}, type {type(val)}: {exc!r}"
)
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.ctx
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)
result = getter.ctx.get_var(key)
if result is Empty:
return self.safe_get("fallback", obj, getter=getter)
return ExpOp(key).eval(getter.ctx, *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
class SetVarHandler(BaseHandler):
input_validators = (Arg.string("lhs"), Arg.present("rhs"))
def handle(self, obj, getter):
key, val = self.safe_get_all(obj, getter)
getter.ctx.set_var(key, val)
return val
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(),
"set_var": SetVarHandler(),
}
BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS
+112
View File
@@ -0,0 +1,112 @@
class ParseError(Exception):
pass
# Pratt parsing referenced from https://github.com/andychu/pratt-parsing-demo
class ParserSpec:
@staticmethod
def null_error(p, token, bp):
raise ParseError(f"{token!r} cannot be used in prefix position")
@staticmethod
def left_error(p, token, bp):
raise ParseError(f"{token!r} cannot be used in infix position")
class LeftInfo:
def __init__(self, led=None, lbp=0, rbp=0):
self.led, self.lbp, self.rbp = led or ParserSpec.left_error, lbp, rbp
class NullInfo:
def __init__(self, nud=None, bp=0):
self.nud, self.bp = nud or ParserSpec.null_error, bp
def __init__(self):
self.null_lookup = {}
self.left_lookup = {}
def add_null(self, bp, nud, tokens):
for token in tokens:
self.null_lookup[token] = self.NullInfo(nud, bp)
if token not in self.left_lookup:
self.left_lookup[token] = self.LeftInfo()
def add_led(self, lbp, rbp, led, tokens):
for token in tokens:
self.left_lookup[token] = self.LeftInfo(led, lbp, rbp)
if token not in self.null_lookup:
self.null_lookup[token] = self.NullInfo(self.null_error)
def add_left(self, bp, led, tokens):
return self.add_led(bp, bp, led, tokens)
def add_leftright(self, bp, led, tokens):
return self.add_led(bp, bp - 1, led, tokens)
def lookup(self, token, is_left):
result = (self.left_lookup if is_left else self.null_lookup).get(token)
if result is None:
raise ParseError(f"Unexpected token {token!r}")
return result
@staticmethod
def get_type(token):
if isinstance(token, (int, float)):
return "number"
if isinstance(token, str) and token.isidentifier():
return "op"
return token
class Parser:
def __init__(self, spec, lexer):
self.spec = spec
self.lexer = lexer
self.token = None
self.token_type = None
self.pos = -1
def advance(self):
if self.lexer is None:
self.token_type = self.token = None
return None
try:
self.token = next(self.lexer)
self.token_type = self.spec.get_type(self.token)
self.pos += 1
except StopIteration:
self.token = self.token_type = self.lexer = None
return self.token
def expect(self, val):
if val is not None and (self.lexer is None or self.token != val):
raise ParseError(f"expected {val!r}, got {self.token!r}")
return self.advance()
def parse_until(self, rbp):
if self.lexer is None:
raise ParseError("unexpected end of input")
spec = self.spec
token, token_type = self.token, self.token_type
self.advance()
ni = spec.lookup(token_type, False)
node = ni.nud(self, token, ni.bp)
while self.lexer:
token, token_type = self.token, self.token_type
li = spec.lookup(token_type, True)
if rbp >= li.lbp:
break
self.advance()
node = li.led(self, token, node, li.rbp)
return node
def go(self):
self.advance()
try:
result = self.parse_until(0)
except ParseError as exc:
raise ParseError(
f"pos {self.pos} at token {self.token!r}: parse error: {exc}"
) from None
if self.lexer:
raise ParseError(f"pos {self.pos}: unexpected end of input")
return result
+213
View File
@@ -0,0 +1,213 @@
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):
value = handlers.get_var(self)
if value is Empty:
raise KeyError(f"No handler for op/var {self}")
return value
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:
def __init__(self, obj, ctx, *args, **kwargs):
self.obj = obj
self.ctx = ctx
self.args = args
self.kwargs = kwargs
def __call__(self, k, *, default=Empty):
obj = self.obj
result = (
obj.kwargs.get_eval(k, self.ctx, *self.args, default=default, **self.kwargs)
if isinstance(k, str)
else obj.args.get_eval(k, self.ctx, *self.args, **self.kwargs)
)
if result is Empty:
raise KeyError(f"Unknown key {k!r}")
return result
class 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_handler(self.name)
if handler is Empty:
raise KeyError(f"No handler for op: {self.name!r}")
return handler(
self, getter=ExprGetter(self, handlers, *args, **kwargs), **kwargs
)
def clone(self):
return self.__class__(self.name, self.args.clone(), self.kwargs.clone())
def pretty_string(self, depth=0):
pad = " " * (depth + 1) * 2
kwargs_str = f", {self.kwargs.pretty_string(depth + 1)}" if self.kwargs else ""
return f"<FUNAP {self.name}\n{pad}{self.args.pretty_string(depth + 1)}{kwargs_str}\n{pad[:-2]}>"
def __repr__(self):
kwargs_str = f", {self.kwargs}" if self.kwargs else ""
return f"<FUNAP:{self.name}{self.args}{kwargs_str}>"
class ExpBoundFunAp(ExpFunAp):
__slots__ = ("fun",)
def __init__(self, name, fun, args, kwargs):
super().__init__(name, args, kwargs)
self.fun = fun
def eval(self, handlers, *args, **kwargs):
def get_evaled(k, default=None):
return (
self.kwargs.get_eval(k, handlers, *args, default=default, **kwargs)
if isinstance(k, str)
else self.args.get_eval(k, handlers, *args, **kwargs)
)
return self.fun(self.name, self.args, *args, getter=get_evaled, **kwargs)
__all__ = (
"ExpBase",
"ExpOp",
"ExpBinOp",
"ExpSym",
"ExpTuple",
"ExpKV",
"ExpDict",
"ExpFunAp",
"ExpBoundFunAp",
)
+36
View File
@@ -0,0 +1,36 @@
import itertools
try:
import torch
except ImportError:
# To facilitate testing.
class torch:
class Tensor:
pass
class WrapGenerator:
def __init__(self, g):
self.g = g
self._value = None
self.ready = False
@property
def value(self):
if not self.ready:
raise ValueError("Value not ready")
return self._value
def __iter__(self):
self._value = yield from self.g
self.ready = True
return self._value
def split_iterable(seq, pred):
it = iter(seq)
while True:
toks = tuple(itertools.takewhile(pred, it))
if toks == ():
break
yield toks
+200
View File
@@ -0,0 +1,200 @@
import functools
from ..latent import ImageBatch
from .util import torch
from .types import Empty
class Arg:
__slots__ = ("name", "default", "validator")
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):
if value is Empty:
if self.default is 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 image(cls, name):
return cls(name, validator=ValidateArg.validate_image)
@classmethod
def numeric(cls, name, default=Empty):
return cls(name, default=default, validator=ValidateArg.validate_numeric)
@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_image(idx, val):
if not isinstance(val, ImageBatch):
raise ValidateError(
f"Expected PIL Image argument at {idx}, got {type(val)}"
)
return val
@staticmethod
def validate_sequence(idx, val, *, item_validator=None):
if not isinstance(val, (list, tuple)):
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
+737
View File
@@ -0,0 +1,737 @@
import os
import torch
import numpy as np
import PIL.Image as PILImage
from . import expression as expr
from . import latent
from .external import MODULES as EXT
from .utils import scale_noise, resolve_value
from .latent import OCSTAESD, ImageBatch
ALLOW_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_UNSAFE_EXPRESSIONS") is not None
ALLOW_ALL_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_ALL_UNSAFE") is not None
EXT_BLEH = EXT.get("bleh")
EXT_SONAR = EXT.get("sonar")
EXT_NNLATENTUPSCALE = EXT.get("nnlatentupscale")
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)
amount = (amount,) * len(dim)
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[-2] * scale[0]), int(t.shape[-1] * scale[1]))
print("SCALE", t.shape[-2:], "->", scale)
if not all(v > 0 for v in scale):
raise ValueError(f"Invalid scale: scale values must be > 0, got: {scale!r}")
return latent.scale_samples(t, scale[1], scale[0], mode=mode)
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.ctx
smin, smax, s, sn = (
ctx.get_var(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 ShapeHandler(expr.BaseHandler):
input_validators = (expr.Arg.tensor("tensor"),)
def handle(self, obj, getter):
t = self.safe_get("tensor", obj, getter)
return expr.types.ExpTuple((*t.shape,))
class TAESDDecodeHandler(expr.BaseHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.string("mode", "sd15"),
)
validate_output = expr.Arg.image("output")
def handle(self, obj, getter):
t, mode = self.safe_get_all(obj, getter)
return OCSTAESD.decode(mode, t)
class TAESDEncodeHandler(expr.BaseHandler):
input_validators = (
expr.Arg.image("image"),
expr.Arg.tensor("reference_latent"),
expr.Arg.string("mode", "sd15"),
)
validate_output = expr.Arg.tensor("output")
def handle(self, obj, getter):
imgbatch, ref, mode = self.safe_get_all(obj, getter)
return OCSTAESD.encode(mode, imgbatch, ref)
class ImgShapeHandler(expr.BaseHandler):
input_validators = (expr.Arg.tensor("image"),)
def handle(self, obj, getter):
imgbatch = self.safe_get("image", obj, getter)
if len(imgbatch) == 0:
raise ValueError("Can't get shape of empty image batch")
isz = imgbatch[0].size
return expr.types.ExpTuple((isz[1], isz[0]))
class ImgPILResizeHandler(expr.BaseHandler):
input_validators = (
expr.Arg.image("image"),
expr.Arg.one_of(
"size",
(
expr.ValidateArg.validate_numeric_scalar,
expr.ValidateArg.validate_numscalar_sequence,
),
),
expr.Arg.string("resample_mode", "bicubic"),
expr.Arg.boolean("absolute_scale", False),
)
validate_output = expr.Arg.image("output")
def handle(self, obj, getter):
imgbatch, size, resample_mode, abs_scale = self.safe_get_all(obj, getter)
if not isinstance(size, tuple):
size = (size, size)
if len(size) != 2 or not all(n > 0 for n in size):
raise ValueError(
"Image resize size parameter must be a positive non-zero number or tuple of positive non-zero height, width"
)
try:
resample_mode = PILImage.Resampling[resample_mode.upper()]
except KeyError:
raise ValueError("Bad resample mode")
if len(imgbatch) == 0:
return imgbatch
if abs_scale:
size = tuple(int(v) for v in size)
else:
imgsize = imgbatch[0].size
size = (int(imgsize[1] * size[0]), int(imgsize[0] * size[1]))
new_size = (size[1], size[0])
return ImageBatch(img.resize(new_size, resample_mode) for img in imgbatch)
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()
if EXT_NNLATENTUPSCALE:
from .latent import scale_nnlatentupscale
class ScaleNNLatentUpscaleHandler(expr.BaseHandler):
input_validators = (
expr.Arg.tensor("tensor"),
expr.Arg.string("mode", "sd1"),
expr.Arg.numeric_scalar("scale", 2.0),
)
output_validator = expr.Arg.tensor("output")
def handle(self, obj, getter):
tensor, mode, scale = self.safe_get_all(obj, getter)
if mode not in {"sd1", "sdxl"}:
raise ValueError(
"Bad mode for t_scale_nnlatentupscale: must be either sd15 or sdxl"
)
return scale_nnlatentupscale(mode, tensor, scale)
HANDLERS["t_scale_nnlatentupscale"] = ScaleNNLatentUpscaleHandler()
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(),
"t_shape": ShapeHandler(),
"t_taesd_decode": TAESDDecodeHandler(),
"unsafe_tensor_method": UnsafeTorchTensorMethodHandler(),
"unsafe_torch": UnsafeTorchHandler(),
}
IMAGE_OP_HANDLERS = {
"img_taesd_encode": TAESDEncodeHandler(),
"img_shape": ImgShapeHandler(),
"img_pil_resize": ImgPILResizeHandler(),
}
HANDLERS |= TENSOR_OP_HANDLERS
HANDLERS |= IMAGE_OP_HANDLERS
+21
View File
@@ -0,0 +1,21 @@
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):
MODULES["sonar"] = importlib.import_module("custom_nodes.ComfyUI-sonar").py
with contextlib.suppress(ImportError, NotImplementedError):
MODULES["nnlatentupscale"] = importlib.import_module(
"custom_nodes.ComfyUi_NNLatentUpscale"
)
__all__ = ("MODULES",)
+510
View File
@@ -0,0 +1,510 @@
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 = {}
FILTER_HANDLERS = expr.HandlerContext(
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_up": ss.sigma_up,
"sigma_prev": ss.sigma_prev,
"hist_len": len(ss.hist),
"sigma_min": ms.sigma_min.item(),
"sigma_max": ms.sigma_max.item(),
"step_pct": float(ss.step / ss.total_steps),
"total_steps": ss.total_steps,
"sampling_pct": (999 - ms.timestep(ss.sigma).item()) / 999,
})
if have_current and len(ss.hist) > 0:
fr |= cls.from_mr(ss.hcur)
fr["d"] = ss.d
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(constants=refs, variables={}))
# 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(constants=refs, variables={}))
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,
}
+269
View File
@@ -0,0 +1,269 @@
import numpy as np
import torch
import torch.nn.functional as F
import folder_paths
import latent_preview
from comfy.taesd.taesd import TAESD
from comfy.utils import bislerp
from comfy import latent_formats
from .external import MODULES as EXT
def normalize_to_scale(latent, target_min, target_max, *, dim=(-3, -2, -1)):
min_val, max_val = (
latent.amin(dim=dim, keepdim=True),
latent.amax(dim=dim, keepdim=True),
)
normalized = (latent - min_val).div_(max_val - min_val)
return (
normalized.mul_(target_max - target_min)
.add_(target_min)
.clamp_(target_min, target_max)
)
# 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())
class ImageBatch(tuple):
__slots__ = ()
class OCSTAESD:
latent_formats = {
"sd15": latent_formats.SD15(),
"sdxl": latent_formats.SDXL(),
}
@classmethod
def get_decoder_name(cls, fmt):
return cls.latent_formats[fmt].taesd_decoder_name
@classmethod
def get_encoder_name(cls, fmt):
result = cls.get_decoder_name(fmt)
if not result.endswith("_decoder"):
raise RuntimeError(
f"Could not determine TAESD encoder name from {result!r}"
)
return f"{result[:-7]}encoder"
@classmethod
def get_taesd_path(cls, name):
taesd_path = next(
(
fn
for fn in folder_paths.get_filename_list("vae_approx")
if fn.startswith(name)
),
"",
)
if taesd_path == "":
raise RuntimeError(f"Could not get TAESD path for {name!r}")
return folder_paths.get_full_path("vae_approx", taesd_path)
@classmethod
def decode(cls, fmt, latent):
latent_format = cls.latent_formats[fmt]
# rv = latent_format.process_out(1.0)
filename = cls.get_taesd_path(cls.get_decoder_name(fmt))
model = TAESD(
decoder_path=filename, latent_channels=latent_format.latent_channels
).to(latent.device)
# print("DEC INPUT ORIG", latent.min(), latent.max())
# if torch.any(latent.max() > rv) or torch.any(latent.min() < -rv):
# sv = latent.new((-rv, rv))
# latent = normalize_to_scale(
# latent,
# latent.amin(dim=(-3, -2, -1), keepdim=True).maximum(sv[0]),
# latent.amax(dim=(-3, -2, -1), keepdim=True).minimum(sv[1]),
# dim=(-3, -2, -1),
# )
# print("DEC INPUT", latent.min(), latent.max())
# result = model.decode(latent.clamp(-rv, rv)).movedim(1, 3)
result = model.decode(latent).movedim(1, 3)
# print("DEC RESULT", result.shape, result.isnan().any().item())
return ImageBatch(
latent_preview.preview_to_image(result[batch_idx])
for batch_idx in range(result.shape[0])
)
@staticmethod
def img_to_encoder_input(imgbatch):
return torch.stack(
tuple(
torch.tensor(np.array(img), dtype=torch.float32)
.div_(127)
.sub_(1.0)
.clamp_(-1, 1)
for img in imgbatch
),
dim=0,
).movedim(-1, 1)
@classmethod
def encode(cls, fmt, imgbatch, latent, *, normalize_output=False):
latent_format = cls.latent_formats[fmt]
rv = latent_format.process_out(1.0)
filename = cls.get_taesd_path(cls.get_encoder_name(fmt))
model = TAESD(
encoder_path=filename, latent_channels=latent_format.latent_channels
).to(device=latent.device)
result = model.encode(cls.img_to_encoder_input(imgbatch).to(latent.device))
# print(
# "ENC RESULT ORIG",
# result.min(),
# result.max(),
# )
# if torch.any(result.max() > rv) or torch.any(result.min() < -rv):
# sv = result.new((-rv, rv))
# result = normalize_to_scale(
# result,
# result.amin(dim=(-3, -2, -1), keepdim=True).maximum(sv[0]),
# result.amax(dim=(-3, -2, -1), keepdim=True).minimum(sv[1]),
# dim=(-3, -2, -1),
# )
# print(
# "ENC RESULT",
# result.shape,
# result.isnan().any().item(),
# result.min(),
# result.max(),
# )
return result.to(latent.dtype).clamp(-rv, rv)
if "bleh" in EXT:
scale_samples = EXT["bleh"].latent_utils.scale_samples
UPSCALE_METHODS = EXT["bleh"].latent_utils.UPSCALE_METHODS
else:
UPSCALE_METHODS = ("bicubic", "bislerp", "bilinear", "nearest-exact", "area")
def scale_samples(
samples,
width,
height,
mode="bicubic",
sigma=None, # noqa: ARG001
):
if mode == "bislerp":
return bislerp(samples, width, height)
return F.interpolate(samples, size=(height, width), mode=mode)
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)
if "nnlatentupscale" in EXT:
def scale_nnlatentupscale(
mode,
latent,
scale=2.0,
*,
scale_factor=0.13025,
__nlu_module=EXT["nnlatentupscale"],
):
module = __nlu_module
mode = {"sdxl": "SDXL", "sd1": "SD 1.x"}.get(mode)
if mode is None:
raise ValueError("Bad mode")
node = module.NNLatentUpscale()
model = module.latent_resizer.LatentResizer.load_model(
node.weight_path[mode], latent.device, latent.dtype
).to(device=latent.device)
result = (
model(scale_factor * latent, scale=scale).to(
dtype=latent.dtype, device=latent.device
)
/ scale_factor
)
del model
return result
+264
View File
@@ -0,0 +1,264 @@
from collections import namedtuple
import torch
import comfy
from comfy.k_diffusion.sampling import to_d
from . import filtering
from .utils import fallback
class History:
def __init__(self, size):
self.history = []
self.size = size
def __len__(self):
return len(self.history)
def __getitem__(self, k):
return self.history[k]
def push(self, val):
if len(self.history) >= self.size:
self.history = self.history[-(self.size - 1) :]
self.history.append(val)
def reset(self):
self.history = []
def clone(self):
obj = self.__new__(self.__class__)
obj.__init__(self.size)
obj.history = self.history.copy()
return obj
class ModelResult:
def __init__(
self,
call_idx,
sigma,
x,
denoised,
**kwargs,
):
self.call_idx = call_idx
self.sigma = sigma
self.x = x
self.denoised = denoised
for k in ("denoised_uncond", "denoised_cond", "tangents", "jdenoised"):
setattr(self, k, kwargs.pop(k, None))
if len(kwargs) != 0:
raise ValueError(f"Unexpected keyword arguments: {tuple(kwargs.keys())}")
def to_d(
self,
/,
x=None,
sigma=None,
denoised=None,
denoised_uncond=None,
alt_cfgpp_scale=0,
cfgpp=False,
):
x = fallback(x, self.x)
sigma = fallback(sigma, self.sigma)
denoised = fallback(denoised, self.denoised)
denoised_uncond = fallback(denoised_uncond, self.denoised_uncond)
if alt_cfgpp_scale != 0:
x = x - denoised * alt_cfgpp_scale + denoised_uncond * alt_cfgpp_scale
return to_d(x, sigma, denoised if not cfgpp else denoised_uncond)
@property
def d(self):
return self.to_d()
def clone(self, deep=False):
obj = self.__new__(self.__class__)
for k in (
"denoised",
"call_idx",
"sigma",
"x",
"denoised_uncond",
"denoised_cond",
"tangents",
"jdenoised",
):
val = getattr(self, k)
if deep and isinstance(val, torch.Tensor):
val = val.copy()
setattr(obj, k, val)
return obj
ModelCallCacheConfig = namedtuple(
"ModelCallCacheConfig", ("size", "max_use", "threshold"), defaults=(0, 1000000, 1)
)
class ModelCallCache:
def __init__(
self,
model,
x,
s_in,
extra_args,
*,
cache=None,
filter=None,
):
self.cache = ModelCallCacheConfig(**fallback(cache, {}))
filtargs = fallback(filter, {}).copy()
self.filters = {}
for key in ("input", "denoised", "jdenoised", "cond", "uncond", "x"):
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", "x"):
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
+533 -79
View File
@@ -1,14 +1,27 @@
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"
DESCRIPTION = "Overly Complicated Sampling main sampler node. Can be connected to a SamplerCustom or other sampler node that supports a SAMPLER input."
OUTPUT_TOOLTIPS = (
"SAMPLER that can be connected to a SamplerCustom or other sampler node that supports a SAMPLER input.",
)
FUNCTION = "go"
@@ -16,34 +29,28 @@ class ComposableSampler:
def INPUT_TYPES(cls):
return {
"required": {
"s_noise": (
"FLOAT",
"groups": (
"OCS_GROUPS",
{
"default": 1.0,
"min": -100.0,
"max": 100.0,
"step": 0.01,
"round": False,
"tooltip": "Connect OCS substep groups here which are output from the OCS Group node."
},
),
"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",),
},
"optional": {
"merge_sampler_opt": ("STEP_SAMPLER_CHAIN",),
"params_opt": (
"OCS_PARAMS",
{
"tooltip": "Optionally connect parameters like custom noise here. Output from the OCS Param or OCS MultiParam nodes.",
},
),
"parameters": (
"STRING",
{"default": "", "multiline": True, "dynamicPrompts": False},
{
"default": DEFAULT_YAML_PARAMS,
"multiline": True,
"dynamicPrompts": False,
"tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.",
},
),
},
}
@@ -51,41 +58,35 @@ 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"
DESCRIPTION = "Over Complicated Sampling group definition node."
OUTPUT_TOOLTIPS = (
"This output can be connect to another OCS Group node or an OCS Sampler node.",
)
FUNCTION = "go"
@@ -93,52 +94,505 @@ class ComposableStepSampler:
def INPUT_TYPES(cls):
return {
"required": {
"s_noise": (
"FLOAT",
"merge_method": (
tuple(MERGE_SUBSTEPS_CLASSES.keys()),
{
"default": 1.0,
"min": -100.0,
"max": 100.0,
"step": 0.01,
"round": False,
"tooltip": "The merge method determines how multiple substeps are combined together during sampling.",
},
),
"eta": (
"FLOAT",
"time_mode": (
("step", "step_pct", "sigma"),
{
"default": 1.0,
"min": -100.0,
"max": 100.0,
"step": 0.01,
"round": False,
"tooltip": "The time mode controls how the time_start and time_end parameters are interpreted. The default of step is generally easiest to use.",
},
),
"time_start": (
"FLOAT",
{
"default": 0,
"min": 0.0,
"step": 0.1,
"round'": False,
"tooltip": "The start time this group will be active (inclusive).",
},
),
"time_end": (
"FLOAT",
{
"default": 999,
"min": 0.0,
"step": 0.1,
"round'": False,
"tooltip": "The group will become inactive when the current time is GREATER than the specified end time.",
},
),
"substeps": (
"OCS_SUBSTEPS",
{
"tooltip": "Connect output from an OCS Substeps node here.",
},
),
"substeps": ("INT", {"default": 1, "min": 1, "max": 1000}),
"step_method": (tuple(STEP_SAMPLERS.keys()),),
},
"optional": {
"step_sampler_opt": ("STEP_SAMPLER_CHAIN",),
"custom_noise_opt": ("SONAR_CUSTOM_NOISE",),
"groups_opt": (
"OCS_GROUPS",
{
"tooltip": "You may optionally connect the output from another OCS Group node here. Only one group per step is used, matching (based on time or other constraints) starts with the OCS Group node furthest from the OCS Sampler.",
},
),
"params_opt": (
"OCS_PARAMS",
{
"tooltip": "Optionally connect parameters like custom noise here. Output from the OCS Param or OCS MultiParam nodes.",
},
),
"parameters": (
"STRING",
{"default": "", "multiline": True, "dynamicPrompts": False},
{
"default": DEFAULT_YAML_PARAMS,
"multiline": True,
"dynamicPrompts": False,
"tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.",
},
),
},
}
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"
DESCRIPTION = "Overly Complicated Sampling substeps definition node. Used to define a sampler type and other sampler-specific parameters."
OUTPUT_TOOLTIPS = (
"This output can be connected to another OCS Substeps node or an OCS Group node.",
)
FUNCTION = "go"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"substeps": (
"INT",
{
"default": 1,
"min": 1,
"max": 1000,
"tooltip": "Number of substeps to use for each step, in other words (depending on the OCS Group merge strategy) it may split a step into multiple smaller steps.",
},
),
"step_method": (
tuple(STEP_SAMPLERS.keys()),
{
"tooltip": "In other words, the sampler.",
},
),
},
"optional": {
"substeps_opt": (
"OCS_SUBSTEPS",
{
"tooltip": "Optionally connect another OCS Substeps node here. Substeps will run in order, starting from the OCS Substeps node FURTHEST from the OCS Group node.",
},
),
"params_opt": (
"OCS_PARAMS",
{
"tooltip": "Optionally connect parameters like custom noise here. Output from the OCS Param or OCS MultiParam nodes.",
},
),
"parameters": (
"STRING",
{
"default": DEFAULT_YAML_PARAMS,
"multiline": True,
"dynamicPrompts": False,
"tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.",
},
),
},
}
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"
DESCRIPTION = "Overly Complicated Sampling parameter definition node. Used to set parameters like custom noise types that require an input."
OUTPUT_TYPES = (
"Can be connected to another OCS Param or OCS MultiParam node or any other OCS node that takes OCS_PARAMS as an input.",
)
FUNCTION = "go"
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()),
{
"tooltip": "Used to set the type of custom parameter.",
},
),
"value": (
cls.WC,
{
"tooltip": "Connect the type of value expected by the key. Allows connecting output from any type of node HOWEVER if it is the wrong type expected by the key you will get an error when you run the workflow.",
},
),
},
"optional": {
"params_opt": (
"OCS_PARAMS",
{
"tooltip": "You may optionally connect the output from other OCS Param or OCS MultiParam nodes here to set multiple parameters.",
},
),
"parameters": (
"STRING",
{
"default": "# Additional YAML or JSON parameters\n",
"multiline": True,
"dynamicPrompts": False,
"tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.",
},
),
},
}
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"
DESCRIPTION = "Overly Complicated Sampling parameter definition node. Used to set parameters like custom noise types that require an input. Like the OCS Param node but allows setting multiple parameters at the same time."
OUTPUT_TYPES = (
"Can be connected to another OCS Param or OCS MultiParam node or any other OCS node that takes OCS_PARAMS as an input.",
)
FUNCTION = "go"
PARAM_COUNT = 5
@classmethod
def INPUT_TYPES(cls):
param_keys = (
("", *ParamNode.OCS_PARAM_TYPES.keys()),
{
"tooltip": "Used to set the type of custom parameter.",
},
)
return {
"required": {
f"key_{idx}": param_keys for idx in range(1, cls.PARAM_COUNT + 1)
},
"optional": {
"params_opt": (
"OCS_PARAMS",
{
"tooltip": "You may optionally connect the output from other OCS MultiParam or OCS Param nodes here to set multiple parameters.",
},
),
"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,
"tooltip": "The text parameter block allows setting custom parameters using YAML (recommended) or JSON. Optional, may be left blank.",
},
),
}
| {
f"value_opt_{idx}": (
ParamNode.WC,
{
"tooltip": "Connect the type of value expected by the corresponding key. Allows connecting output from any type of node HOWEVER if it is the wrong type expected by the corresponding key you will get an error when you run the workflow.",
},
)
for idx in range(1, cls.PARAM_COUNT + 1)
},
}
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"
DESCRIPTION = "Overly Complicated Sampling simple Restart schedule node. Allows generating a Restart sampling schedule based on a text definition."
OUTPUT_TYPES = (
"Can be connected to an OCS Sampler or RestartSampler node. Do not connect directly to a sampler that doesn't have built-in support for Restart schedules.",
)
FUNCTION = "go"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"sigmas": (
"SIGMAS",
{
"tooltip": "Connect the output from another scheduler node (i.e. BasicScheduler) here.",
},
),
"start_step": (
"INT",
{
"min": 0,
"default": 0,
"tooltip": "Step the restart schedule definition starts applying. Zero-based.",
},
),
},
"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,
"tooltip": "Define a schedule here using YAML (recommended) or JSON.",
},
),
},
}
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"
DESCRIPTION = "Allows forcing a model's maximum and minumum sigmas to a specified value. You generally do NOT want to connect this to a sampler node. Connect it to a scheduler node (i.e. BasicScheduler) instead."
OUTPUT_TOOLTIPS = (
"Patched model. Can be connected to a scheduler node (i.e. BasicScheduler). Generally NOT recommended to connect to an actual sampler.",
)
FUNCTION = "go"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": (
"MODEL",
{
"tooltip": "Model to patch with the min/max sigmas.",
},
),
"mode": (
("recalculate", "simple_multiply"),
{
"tooltip": "Mode use for setting sigmas in the patched model. Recalculate should generally be more accurate.",
},
),
"sigma_max": (
"FLOAT",
{
"default": -1.0,
"min": -10000.0,
"max": 10000.0,
"step": 0.01,
"round": False,
"tooltip": "You can set the maximum sigma here. If you use a negative value, it will be interpreted as the absolute value for the max sigma. If you use a positive value it will be interpreted as a percentage (where 1.0 signified 100%). Schedules generated with the patched model should start from sigma_max (or close to it).",
},
),
"fake_sigma_min": (
"FLOAT",
{
"default": 0.0,
"min": 0.0,
"max": 1000.0,
"step": 0.01,
"round": False,
"tooltip": "You can set the minimum sigma here. Disabled if set to 0. If you use a negative value, it will be interpreted as the absolute value for the max sigma. If you use a positive value it will be interpreted as a percentage (where 1.0 signified 100%). Schedules generated with the patched model should end with [sigma_min, 0]. NOTE: May not work with some schedulers. I recommend leaving this at 0 unless you know you need it (and even then it may not work).",
},
),
}
}
def go(self, model, mode="recalculate", sigma_max=-1.0, fake_sigma_min=0.0):
if sigma_max == 0:
raise ValueError("ModelSetMaxSigma: Invalid sigma_max value")
if mode not in ("recalculate", "simple_multiply"):
raise ValueError("ModelSetMaxSigma: Invalid mode value")
orig_ms = model.get_model_object("model_sampling")
model = model.clone()
orig_max_sigma, orig_min_sigma = (
orig_ms.sigma_max.item(),
orig_ms.sigma_min.item(),
)
max_multiplier = abs(sigma_max) if sigma_max < 0 else sigma_max / orig_max_sigma
if max_multiplier == 1:
return (model,)
mcfg = model.get_model_object("model_config")
orig_sigmas = orig_ms.sigmas
fake_sigma_min = orig_sigmas.new_full((1,), fake_sigma_min)
class NewModelSampling(orig_ms.__class__):
if fake_sigma_min != 0:
@property
def sigma_min(self):
return fake_sigma_min
ms = NewModelSampling(mcfg)
if mode == "simple_multiply":
ms.set_sigmas(orig_sigmas * max_multiplier)
else:
ss = getattr(mcfg, "sampling_setting", None) or {}
if ss.get("beta_schedule", "linear") != "linear":
raise NotImplementedError(
"ModelSetMaxSigma: Can only handle linear beta schedules in reschedule mode"
)
ms.set_sigmas((orig_sigmas**2 * max_multiplier**2) ** 0.5)
new_max_sigma, new_min_sigma = ms.sigma_max.item(), ms.sigma_min.item()
if new_min_sigma >= new_max_sigma:
raise ValueError(
"ModelSetMaxSigma: Invalid fake_min_sigma value, result max <= min"
)
model.add_object_patch("model_sampling", ms)
print(
f"ModelSetMaxSigma: Set model sigmas({mode}): old_max={orig_max_sigma:.04}, old_min={orig_min_sigma:.03}, new_max={new_max_sigma:.04}, new_min={new_min_sigma:.03}"
)
return (model,)
__all__ = (
"SamplerNode",
"GroupNode",
"SubstepsNode",
"ParamNode",
"MultiParamNode",
"ModelSetMaxSigmaNode",
)
+231
View File
@@ -0,0 +1,231 @@
import gc
import random
import scipy
import torch
from .filtering import Filter, make_filter
from .utils import scale_noise, fallback
class ImmiscibleNoise(Filter):
name = "immiscible"
uses_ref = True
default_options = Filter.default_options | {
"size": 0,
"batching": "channel",
"maximize": False,
}
def __call__(self, noise_sampler, x_ref, *, refs=None):
if not self.check_applies(refs):
return noise_sampler()
return self.apply(
torch.cat(tuple(noise_sampler() for _ in range(self.size)))
if self.size > 0
else noise_sampler(),
default_ref=x_ref,
refs=refs,
output_shape=x_ref.shape,
)
def filter(self, latent, ref_latent, *, refs, output_shape):
if self.size == 0:
return latent
return self.unbatch(
self.immiscible(self.batch(latent), self.batch(ref_latent)), output_shape
)
def batch(self, latent):
if self.batching == "batch":
return latent
sz = latent.shape
if latent.ndim != 4:
raise ValueError("Both latent and reference must be four-dimensional")
if self.batching == "channel":
return latent.view(sz[0] * sz[1], *sz[2:])
if self.batching == "row":
return latent.view(sz[0] * sz[1] * sz[2], sz[3])
if self.batching == "column":
return latent.permute(0, 1, 3, 2).reshape(sz[0] * sz[1] * sz[3], sz[2])
raise ValueError("Bad Immmiscible noise batching type")
def unbatch(self, latent, sz):
if self.batching == "column":
return latent.view(*sz[:2], sz[3], sz[2]).permute(0, 1, 3, 2)
return latent.view(*sz)
# Based on implementation from https://github.com/kohya-ss/sd-scripts/pull/1395
# Idea from https://github.com/Clybius
def immiscible(self, latent, ref_latent):
# "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303
# Minimize latent-noise pairs over a batch
n = latent.shape[0]
ref_latent_expanded = (
ref_latent.half().unsqueeze(1).expand(-1, n, *ref_latent.shape[1:])
)
latent_expanded = (
latent.half().unsqueeze(0).expand(ref_latent.shape[0], *latent.shape)
)
dist = (ref_latent_expanded - latent_expanded) ** 2
dist = dist.mean(list(range(2, dist.dim()))).cpu()
try:
assign_mat = scipy.optimize.linear_sum_assignment(
dist, maximize=self.maximize
)
except ValueError as _exc:
# print("\nImmiscible: Failed optimization, skipping")
return latent[: ref_latent.shape[0]]
# print("IMM IDX", assign_mat[1])
return latent[assign_mat[1]]
class NoiseSamplerCache:
def __init__(
self,
x,
seed,
min_sigma,
max_sigma,
*,
normalize_noise=True,
cpu_noise=True,
batch_size=32,
caching=False,
cache_reset_interval=9999,
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
+94
View File
@@ -0,0 +1,94 @@
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
def __repr__(self):
return f"<Restart: s_noise={self.s_noise:.04}, immiscible={self.immiscible}>"
@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)
+81 -47
View File
@@ -2,9 +2,23 @@ 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(constants=ss.refs)
# 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 +28,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 +43,84 @@ 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()
nsc.update_x(x)
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)
nsc.update_x(x)
# print(
# f"STEP {step + 1:>3}: {ss.sigma.item():.03} -> {ss.sigma_next.item():.03} || up={ss.sigma_up.item():.03}, down={ss.sigma_down.item():.03}"
# )
ss.model.reset_cache()
nsc.update_x(x)
merge_sampler = find_merge_sampler(merge_samplers, ss)
if merge_sampler is None:
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
+2214
View File
File diff suppressed because it is too large Load Diff
+463 -197
View File
@@ -1,214 +1,379 @@
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.preview_mode = options.pop("preview_mode", "denoised")
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
def callback(self, *, ss=None, mr=None, preview_mode=None):
ss = fallback(ss, self.ss)
preview_mode = fallback(preview_mode, self.preview_mode)
return ss.callback(hi=mr, preview_mode=preview_mode)
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)
self.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)
self.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)
self.schedule_multiplier = schedule_multiplier
name = "divide"
def __init__(self, ss, group, **kwargs):
super().__init__(ss, group, **kwargs)
self.schedule_multiplier = self.options.pop("schedule_multiplier", 4)
def make_schedule(self, ss):
max_steps = len(self.ss.sigmas) - 1
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 +384,143 @@ 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:
self.callback(ss=subss)
sr = self.simple_substep(x, ssampler, ss=subss)
x = sr.x
noise_strength = sr.noise_scale
if noise_strength != 0 and subss.sigma_next != 0:
x = sr.noise_x(ss=subss)
substep += 1
pbar.update(1)
pbar.update(0)
return x
class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
name = "overshoot"
def __init__(
self,
ss,
group,
**kwargs,
):
super().__init__(ss, group, **kwargs)
self.overshoot_expand_steps = self.options.pop("overshoot_expand_steps", 1)
restart = self.options.pop("restart", {})
self.restart = Restart(
s_noise=restart.get("s_noise", 1.0),
custom_noise=self.options.pop("restart_custom_noise", None),
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)
self.callback(ss=subss)
sr = self.simple_substep(x, ssampler, ss=subss)
x = sr.x
noise_strength = sr.noise_scale
if noise_strength != 0 and subss.sigma_next != 0:
x = sr.noise_x(ss=subss)
substep += 1
pbar.update(1)
last_down = subss.sigma_next.item()
if subss.idx + substep >= max_idx:
break
if subss.idx >= max_idx:
break
if last_down is not None and last_down < ss.sigma_next:
restart_ns = self.restart.get_noise_sampler(ss.noise)
x += ss.noise.scale_noise(
restart_ns(refs=ss.refs),
self.restart.get_noise_scale(last_down, ss.sigma_next),
)
pbar.update(0)
return x
MERGE_SUBSTEPS_CLASSES = {
"default (simple)": SimpleSubstepsSampler,
"normal": NormalMergeSubstepsSampler,
"divide": DivideMergeSubstepsSampler,
"average": AverageMergeSubstepsSampler,
"sample": SampleMergeSubstepsSampler,
"sample_uncached": SampleUncachedMergeSubstepsSampler,
"overshoot": OvershootMergeSubstepsSampler,
# "average": AverageMergeSubstepsSampler,
# "sample": SampleMergeSubstepsSampler,
# "sample_uncached": SampleUncachedMergeSubstepsSampler,
"simple": SimpleSubstepsSampler,
}
-624
View File
@@ -1,624 +0,0 @@
import math
import torch
from comfy.k_diffusion.sampling import (
get_ancestral_step,
to_d,
)
from .res_support import _de_second_order
from .utils import find_first_unsorted
class SingleStepSampler:
name = None
def __init__(
self,
*,
noise_sampler=None,
substeps=1,
s_noise=1.0,
eta=1.0,
dyn_eta_start=None,
dyn_eta_end=None,
weight=1.0,
**kwargs,
):
self.s_noise = s_noise
self.eta = eta
self.dyn_eta_start = dyn_eta_start
self.dyn_eta_end = dyn_eta_end
self.noise_sampler = noise_sampler
self.weight = weight
self.substeps = substeps
self.kwargs = kwargs
def step(self, x, ss):
raise NotImplementedError
# Euler - based on original ComfyUI implementation
def euler_step(self, x, ss):
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
d = to_d(x, ss.sigma, ss.denoised)
dt = sigma_down - ss.sigma
return x + d * dt, sigma_up
def __str__(self):
return f"<SS({self.name}): s_noise={self.s_noise}, eta={self.eta}>"
def get_dyn_value(self, ss, start, end):
if None in (start, end):
return 1.0
if start == end:
return start
main_idx = getattr(ss, "main_idx", ss.idx)
main_sigmas = getattr(ss, "main_sigmas", ss.sigmas)
step_pct = main_idx / (len(main_sigmas) - 1)
dd_diff = end - start
return start + dd_diff * step_pct
def get_dyn_eta(self, ss):
return self.eta * self.get_dyn_value(ss, self.dyn_eta_start, self.dyn_eta_end)
class ReversibleSingleStepSampler(SingleStepSampler):
def __init__(self, *, reta=1.0, dyn_reta_start=None, dyn_reta_end=None, **kwargs):
super().__init__(**kwargs)
self.reta = reta
self.dyn_reta_start = dyn_reta_start
self.dyn_reta_end = dyn_reta_end
def get_dyn_reta(self, ss):
return self.reta * self.get_dyn_value(
ss, self.dyn_reta_start, self.dyn_reta_end
)
class EulerStep(SingleStepSampler):
name = "euler"
step = SingleStepSampler.euler_step
class DPMPPStepBase(SingleStepSampler):
@staticmethod
def sigma_fn(t):
return t.neg().exp()
@staticmethod
def t_fn(t):
return t.log().neg()
class DPMPP2MStep(DPMPPStepBase):
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
t, t_next = self.t_fn(ss.sigma), self.t_fn(ss.sigma_next)
h = t_next - t
st, st_next = self.sigma_fn(t), self.sigma_fn(t_next)
if len(ss.dhist) == 0 or ss.sigma_prev is None:
return (st_next / st) * x - (-h).expm1() * ss.denoised, 0.0
h_last = t - self.t_fn(ss.sigma_prev)
r = h_last / h
denoised, old_denoised = ss.denoised, ss.dhist[-1]
denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised
return (st_next / st) * x - (-h).expm1() * denoised_d, 0.0
class DPMPP2MSDEStep(SingleStepSampler):
name = "dpmpp_2m_sde"
def __init__(self, *, solver_type="midpoint", **kwargs):
super().__init__(**kwargs)
self.solver_type = solver_type
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
denoised = ss.denoised
if ss.sigma_next == 0:
return denoised, None
# DPM-Solver++(2M) SDE
t, s = -ss.sigma.log(), -ss.sigma_next.log()
h = s - t
eta_h = self.get_dyn_eta(ss) * h
x = (
ss.sigma_next / ss.sigma * (-eta_h).exp() * x
+ (-h - eta_h).expm1().neg() * denoised
)
noise_strength = ss.sigma_next * (-2 * eta_h).expm1().neg().sqrt()
if len(ss.dhist) == 0 or ss.sigma_prev is None:
return x, noise_strength
h_last = (-ss.sigma.log()) - (-ss.sigma_prev.log())
r = h_last / h
old_denoised = ss.dhist[-1]
if self.solver_type == "heun":
x = x + (
((-h - eta_h).expm1().neg() / (-h - eta_h) + 1)
* (1 / r)
* (denoised - old_denoised)
)
elif self.solver_type == "midpoint":
x = x + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * (
denoised - old_denoised
)
return x, noise_strength
class DPMPP3MSDEStep(SingleStepSampler):
name = "dpmpp_3m_sde"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
denoised = ss.denoised
if ss.sigma_next == 0:
return denoised, 0
t, s = -ss.sigma.log(), -ss.sigma_next.log()
h = s - t
eta = self.get_dyn_eta(ss)
h_eta = h * (eta + 1)
x = torch.exp(-h_eta) * x + (-h_eta).expm1().neg() * denoised
noise_strength = ss.sigma_next * (-2 * h * eta).expm1().neg().sqrt()
if len(ss.dhist) == 0 or ss.sigma_prev is None:
return x, noise_strength
h_1 = (-ss.sigma.log()) - (-ss.sigma_prev.log())
denoised_1 = ss.dhist[-1]
if len(ss.dhist) == 1:
r = h_1 / h
d = (denoised - denoised_1) / r
phi_2 = h_eta.neg().expm1() / h_eta + 1
x = x + phi_2 * d
else:
h_2 = (-ss.sigma_prev.log()) - (-ss.sigmas[ss.idx - 2].log())
denoised_2 = ss.dhist[-2]
r0 = h_1 / h
r1 = h_2 / h
d1_0 = (denoised - denoised_1) / r0
d1_1 = (denoised_1 - denoised_2) / r1
d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1)
d2 = (d1_0 - d1_1) / (r0 + r1)
phi_2 = h_eta.neg().expm1() / h_eta + 1
phi_3 = phi_2 / h_eta - 0.5
x = x + phi_2 * d1 - phi_3 * d2
return x, noise_strength
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class ReversibleHeunStep(ReversibleSingleStepSampler):
name = "reversible_heun"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(
self.get_dyn_reta(ss)
)
dt = sigma_down - ss.sigma
dt_reversible = sigma_down_reversible - ss.sigma
# Calculate the derivative using the model
d = to_d(x, ss.sigma, ss.denoised)
# Predict the sample at the next sigma using Euler step
x_pred = x + d * dt
# Denoised sample at the next sigma
denoised_next = ss.model(x_pred, sigma_down, model_call_idx=1)
# Calculate the derivative at the next sigma
d_next = to_d(x_pred, sigma_down, denoised_next)
# Update the sample using the Reversible Heun formula
x = x + dt * (d + d_next) / 2 - dt_reversible**2 * (d_next - d) / 4
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class ReversibleHeun1SStep(ReversibleSingleStepSampler):
name = "reversible_heun_1s"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
# Reversible Heun-inspired update (first-order)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(
self.get_dyn_reta(ss)
)
sigma_i, sigma_i_plus_1 = ss.sigma, sigma_down
dt = sigma_i_plus_1 - sigma_i
dt_reversible = sigma_down_reversible - sigma_i
eff_x = ss.xhist[-1] if len(ss.xhist) else x
# Calculate the derivative using the model
d_i_old = to_d(
eff_x,
sigma_i,
ss.dhist[-1]
if len(ss.dhist)
else ss.model(eff_x, sigma_i, model_call_idx=1),
)
# Predict the sample at the next sigma using Euler step
x_pred = eff_x + d_i_old * dt
# Calculate the derivative at the next sigma
d_i_plus_1 = to_d(x_pred, sigma_i_plus_1, ss.denoised)
# Update the sample using the Reversible Heun formula
x = (
x
+ dt * (d_i_old + d_i_plus_1) / 2
- dt_reversible**2 * (d_i_plus_1 - d_i_old) / 4
)
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class RESStep(SingleStepSampler):
name = "res"
def __init__(self, *, res_simple_phi=False, res_c2=0.5, **kwargs):
super().__init__(**kwargs)
self.simple_phi = res_simple_phi
self.c2 = res_c2
pass
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
eta = self.get_dyn_eta(ss)
sigma_down, sigma_up = ss.get_ancestral_step(eta)
denoised = ss.denoised
lam_next = sigma_down.log().neg() if eta != 0 else ss.sigma_next.log().neg()
lam = ss.sigma.log().neg()
h = lam_next - lam
a2_1, b1, b2 = _de_second_order(
h=h, c2=self.c2, simple_phi_calc=self.simple_phi
)
c2_h = 0.5 * h
x_2 = math.exp(-c2_h) * x + a2_1 * h * denoised
lam_2 = lam + c2_h
sigma_2 = lam_2.neg().exp()
denoised2 = ss.model(x_2, sigma_2, model_call_idx=1)
x = math.exp(-h) * x + h * (b1 * denoised + b2 * denoised2)
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class TrapezoidalStep(SingleStepSampler):
name = "trapezoidal"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
dt = ss.sigma_next - ss.sigma
denoised = ss.denoised
# Calculate the derivative using the model
d_i = to_d(x, ss.sigma, denoised)
# Predict the sample at the next sigma using Euler step
x_pred = x + d_i * dt
# Denoised sample at the next sigma
denoised_next = ss.model(x_pred, ss.sigma_next, model_call_idx=1)
# Calculate the derivative at the next sigma
d_next = to_d(x_pred, ss.sigma_next, denoised_next)
dt_2 = sigma_down - ss.sigma
# Update the sample using the Trapezoidal rule
x = x + dt_2 * (d_i + d_next) / 2
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class BogackiStep(ReversibleSingleStepSampler):
name = "bogacki"
reversible = False
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(
self.get_dyn_reta(ss)
)
sigma, sigma_next = ss.sigma, sigma_down
dt = sigma_next - sigma
dt_reversible = sigma_down_reversible - sigma
denoised = ss.denoised
# Calculate the derivative using the model
d = to_d(x, sigma, denoised)
# Bogacki-Shampine steps
k1 = d * dt
k2 = (
to_d(
x + k1 / 2,
sigma + dt / 2,
ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=1),
)
* dt
)
k3 = (
to_d(
x + 3 * k1 / 4 + k2 / 4,
sigma + 3 * dt / 4,
ss.model(x + 3 * k1 / 4 + k2 / 4, sigma + 3 * dt / 4, model_call_idx=2),
)
* dt
)
# Reversible correction term (inspired by Reversible Heun)
correction = dt_reversible**2 * (k3 - k2) / 6 if self.reversible else 0.0
# Update the sample
x = x + 2 * k1 / 9 + k2 / 3 + 4 * k3 / 9 - correction
return x, sigma_up
class ReversibleBogackiStep(BogackiStep):
name = "reversible_bogacki"
reversible = True
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class RK4Step(SingleStepSampler):
name = "rk4"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
sigma = ss.sigma
# Calculate the derivative using the model
d = to_d(x, sigma, ss.denoised)
dt = sigma_down - sigma
# Runge-Kutta steps
k1 = d * dt
k2 = (
to_d(
x + k1 / 2,
sigma + dt / 2,
ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=1),
)
* dt
)
k3 = (
to_d(
x + k2 / 2,
sigma + dt / 2,
ss.model(x + k2 / 2, sigma + dt / 2, model_call_idx=2),
)
* dt
)
k4 = (
to_d(
x + k3,
sigma + dt,
ss.model(x + k3, sigma + dt, model_call_idx=3),
)
* dt
)
# Update the sample
x = x + (k1 + 2 * k2 + 2 * k3 + k4) / 6
return x, sigma_up
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
class EulerDancingStep(SingleStepSampler):
name = "euler_dancing"
def __init__(
self,
*,
deta=1.0,
ds_noise=1.0,
leap=2,
dyn_deta_start=None,
dyn_deta_end=None,
dyn_deta_mode="lerp",
**kwargs,
):
super().__init__(**kwargs)
self.deta = deta
self.ds_noise = ds_noise
self.leap = leap
self.dyn_deta_start = dyn_deta_start
self.dyn_deta_end = dyn_deta_end
if dyn_deta_mode not in ("lerp", "lerp_alt", "deta"):
raise ValueError("Bad dyn_deta_mode")
self.dyn_deta_mode = dyn_deta_mode
def step(self, x, ss):
eta = self.get_dyn_eta(ss)
leap_sigmas = ss.sigmas[ss.idx :]
leap_sigmas = leap_sigmas[: find_first_unsorted(leap_sigmas)]
zero_idx = (leap_sigmas <= 0).nonzero().flatten()[:1]
max_leap = (zero_idx.item() if len(zero_idx) else len(leap_sigmas)) - 1
is_danceable = max_leap > 1 and ss.sigma_next != 0
curr_leap = max(1, min(self.leap, max_leap))
sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next
print("DANCE", max_leap, curr_leap, sigma_leap, "--", leap_sigmas)
del leap_sigmas
sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta)
d = to_d(x, ss.sigma, ss.denoised)
# Euler method
dt = sigma_down - ss.sigma
x = x + d * dt
if curr_leap == 1:
return x, sigma_up
dance_scale = self.get_dyn_value(ss, self.dyn_deta_start, self.dyn_deta_end)
if not is_danceable or abs(dance_scale) < 1e-04:
return x, sigma_up
sigma_down_normal, sigma_up_normal = get_ancestral_step(
ss.sigma, ss.sigma_next, eta
)
if self.dyn_deta_mode == "lerp":
dt_normal = sigma_down_normal - ss.sigma
x_normal = x + d * dt_normal
else:
x_normal = x
x = x + self.noise_sampler(ss.sigma, sigma_leap) * self.s_noise * sigma_up
sigma_down2, sigma_up2 = get_ancestral_step(
sigma_leap,
ss.sigma_next,
eta=self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale),
)
d_2 = to_d(x, sigma_leap, ss.denoised)
dt_2 = sigma_down2 - sigma_leap
result = x + d_2 * dt_2
noise_diff = sigma_up2 - sigma_up * dance_scale
noise_scale = sigma_up2 + noise_diff * (0.025 * curr_leap)
if self.dyn_deta_mode == "deta" or dance_scale == 1.0:
return result, noise_scale
result = torch.lerp(x_normal, result, dance_scale)
# FIXME: Broken for noise samplers that care about s/sn
return result, noise_scale
class DPMPP2SStep(DPMPPStepBase):
name = "dpmpp_2s"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
t_fn, sigma_fn = self.t_fn, self.sigma_fn
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
# DPM-Solver++(2S)
t, t_next = t_fn(ss.sigma), t_fn(sigma_down)
r = 1 / 2
h = t_next - t
s = t + r * h
x_2 = (sigma_fn(s) / sigma_fn(t)) * x - (-h * r).expm1() * ss.denoised
denoised_2 = ss.model(x_2, sigma_fn(s), model_call_idx=0)
x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_2
return x, sigma_up
class DPMPPSDEStep(DPMPPStepBase):
name = "dpmpp_sde"
def __init__(self, *args, r=1 / 2, **kwargs):
super().__init__(*args, **kwargs)
self.r = r
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
t_fn, sigma_fn = self.t_fn, self.sigma_fn
r, eta, s_noise = self.r, self.get_dyn_eta(ss), self.s_noise
noise_sampler = self.noise_sampler
sigma_down, sigma_up = ss.get_ancestral_step(eta)
# DPM-Solver++
t, t_next = t_fn(ss.sigma), t_fn(ss.sigma_next)
h = t_next - t
s = t + h * r
fac = 1 / (2 * r)
# Step 1
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta)
s_ = t_fn(sd)
x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - (t - s_).expm1() * ss.denoised
x_2 = x_2 + noise_sampler(sigma_fn(t), sigma_fn(s)) * s_noise * su
denoised_2 = ss.model(x_2, sigma_fn(s), model_call_idx=1)
# Step 2
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta)
t_next_ = t_fn(sd)
denoised_d = (1 - fac) * ss.denoised + fac * denoised_2
x = (sigma_fn(t_next_) / sigma_fn(t)) * x - (t - t_next_).expm1() * denoised_d
return x, su
# Based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
# Which was originally written by Katherine Crowson
class TTMJVPStep(SingleStepSampler):
name = "ttm_jvp"
def __init__(self, *args, alternate_phi_2_calc=True, **kwargs):
super().__init__(*args, **kwargs)
self.alternate_phi_2_calc = alternate_phi_2_calc
def step(self, x, ss):
if ss.sigma_next == 0:
return ss.denoised, ss.sigma.new_zeros(1)
eta = self.get_dyn_eta(ss)
sigma_down, sigma_up = ss.get_ancestral_step(eta)
sigma, sigma_next = ss.sigma, ss.sigma_next
# 2nd order truncated Taylor method
t, s = -sigma.log(), -sigma_next.log()
h = s - t
h_eta = h * (eta + 1)
eps = to_d(x, sigma, ss.denoised)
denoised, denoised_prime = ss.model(
x, sigma, tangents=(eps * -sigma, -sigma), model_call_idx=1
)
phi_1 = -torch.expm1(-h_eta)
if self.alternate_phi_2_calc:
phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0
else:
phi_2 = torch.expm1(-h_eta) + h_eta
x = torch.exp(-h_eta) * x + phi_1 * ss.denoised + phi_2 * denoised_prime
if not eta:
return x, ss.sigma.new_zeros(1)
phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * eta))
return x, sigma_next * phi_1_noise
STEP_SAMPLERS = {
"euler": EulerStep,
"dpmpp_sde": DPMPPSDEStep,
"dpmpp_2m": DPMPP2MStep,
"dpmpp_2m_sde": DPMPP2MSDEStep,
"dpmpp_3m_sde": DPMPP3MSDEStep,
"dpmpp_2s": DPMPP2SStep,
"reversible_heun": ReversibleHeunStep,
"reversible_heun_1s": ReversibleHeun1SStep,
"res": RESStep,
"trapezoidal": TrapezoidalStep,
"bogacki": BogackiStep,
"reversible_bogacki": ReversibleBogackiStep,
"rk4": RK4Step,
"euler_dancing": EulerDancingStep,
"ttm_jvp": TTMJVPStep,
}
__all__ = (
"STEP_SAMPLERS",
"EulerStep",
"DPMPP2MStep",
"DPMPP2MSDEStep",
"DPMPP3MSDEStep",
"DPMPP2SStep",
"ReversibleHeunStep",
"ReversibleHeun1SStep",
"RESStep",
"TrapezoidalStep",
"BogackiStep",
"ReversibleBogackiStep",
"EulerDancingStep",
"TTMJVPStep",
)
+150 -122
View File
@@ -2,188 +2,216 @@ import torch
from comfy.k_diffusion.sampling import get_ancestral_step
from .filtering import FilterRefs
from .model import History
from .utils import fallback
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, *, preview_mode="denoised"):
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
if preview_mode == "cond":
preview = fallback(hi.denoised_cond, hi.denoised)
elif preview_mode == "uncond":
preview = fallback(hi.denoised_uncond, hi.denoised)
elif preview_mode == "raw":
preview = hi.x
else:
preview = hi.denoised
return self.callback_({
"x": hi.x,
"i": self.step,
"sigma": hi.sigma,
"sigma_hat": hi.sigma,
"denoised": preview,
})
def reset(self):
self.hist.reset()
self.denoised = None
+121 -10
View File
@@ -1,17 +1,77 @@
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):
from . import latent
# def scale_noise_(
# noise,
# factor=1.0,
# *,
# normalized=True,
# normalize_dims=(-3, -2, -1),
# ):
# if not normalized or noise.numel() == 0:
# return noise.mul_(factor) if factor != 1 else noise
# mean, std = (
# noise.mean(dim=normalize_dims, keepdim=True),
# noise.std(dim=normalize_dims, keepdim=True),
# )
# return latent.normalize_to_scale(
# noise.sub_(mean).div_(std).clamp(-1, 1), -1.0, 1.0, dim=normalize_dims
# ).mul_(factor)
# def scale_noise(
# noise,
# factor=1.0,
# *,
# normalized=True,
# normalize_dims=(-3, -2, -1),
# ):
# if not normalized or noise.numel() == 0:
# return noise * factor if factor != 1 else noise
# mean, std = (
# noise.mean(dim=normalize_dims, keepdim=True),
# noise.std(dim=normalize_dims, keepdim=True),
# )
# return (noise - mean).div_(std).mul_(factor)
def scale_noise(
noise,
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
return noise * factor if factor != 1 else noise
noise = noise / noise.std(dim=normalize_dims, keepdim=True)
return noise.sub_(noise.mean(dim=normalize_dims, keepdim=True)).mul_(factor)
# def scale_noise(
# noise,
# factor=1.0,
# *,
# normalized=True,
# normalize_dims=(-3, -2, -1),
# ):
# if not normalized or noise.numel() == 0:
# return noise.mul_(factor) if factor != 1 else noise
# n = (
# torch.nn.LayerNorm(noise.shape[1:])
# if normalize_dims == (-3, -2, -1)
# else torch.nn.InstanceNorm2d(noise.shape[1])
# ).to(noise)
# return n(noise) * factor
# return latent.normalize_to_scale(
# n(noise).clamp_(-1, 1), -1, 1, dim=normalize_dims
# ).mul_(factor)
def find_first_unsorted(tensor, desc=True):
@@ -20,3 +80,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")