Compare commits
20
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
650467ce97 | ||
|
|
bba5bf25e9 | ||
|
|
3a753c1a8b | ||
|
|
ec7def5723 | ||
|
|
4ec5970128 | ||
|
|
e4b05c506d | ||
|
|
6c4ae67e32 | ||
|
|
cf90ae74e1 | ||
|
|
4a97ad3468 | ||
|
|
c2a93d55cb | ||
|
|
2b2a76bcbe | ||
|
|
29ed97230e | ||
|
|
c3d1149aff | ||
|
|
83460f3b8f | ||
|
|
9dedbeb0b0 | ||
|
|
d25d01542e | ||
|
|
1295521583 | ||
|
|
607868c5c1 | ||
|
|
543f39ebf2 | ||
|
|
8097f26863 |
@@ -25,6 +25,7 @@ composite and otherwise manipulate noise see:
|
||||
* [Advanced Power Noise](docs/advanced_power_noise.md) - examples and descriptions of the advanced power noise node.
|
||||
* [Advanced Noise Nodes](docs/advanced_noise_nodes.md) - examples and descriptions of advanced noise nodes (schedule, composite, etc).
|
||||
* [FreeU Extreme](docs/frux.md) - a build your own FreeU kit that allows advanced filtering, blending, scheduling of effects as well as targetting input and middle blocks.
|
||||
* [Wavelet CFG](docs/waveletcfg.md) - replacement CFG function that lets you set different CFG scales for high/low frequency parts of the latent. You can even do stuff like use a different CFG scale for horizontal versus vertical.
|
||||
|
||||
## Sonar Description
|
||||
|
||||
@@ -154,6 +155,7 @@ My version was initially based on this Sonar sampler implementation for Diffuser
|
||||
* Original `SonarPowerNoise` contributed by [elias-gaeros](https://github.com/elias-gaeros/). Additionally, he provided a lot of guidance with refactoring it to allow separate filtering and other enhancements and answered a multitude of dumb questions. To say those changes are only co-authored is probably giving myself too much credit. Thank you! Your patience and help is very much appreciated.
|
||||
* New 1/f (onef) and power law (white, grey, violet, velvet) noise types referenced from https://github.com/WASasquatch/PowerNoiseSuite
|
||||
* Wavelet noise idea (and some of the default settings) from https://github.com/ClownsharkBatwing/RES4LYF
|
||||
* Pattern break algorithm adapted from https://github.com/Extraltodeus/noise_latent_perlinpinpin
|
||||
|
||||
## Errata
|
||||
|
||||
|
||||
+17
-9
@@ -1,15 +1,23 @@
|
||||
from .py import freeu_extreme, nodes, powernoise, sonar
|
||||
import sys
|
||||
|
||||
from . import py # noqa: F401
|
||||
from .py import nodes, sonar
|
||||
|
||||
|
||||
def blep_init():
|
||||
bi = sys.modules.get("_blepping_integrations", {})
|
||||
if "sonar" in bi:
|
||||
return
|
||||
bi["sonar"] = sys.modules[__name__]
|
||||
sys.modules["_blepping_integrations"] = bi
|
||||
|
||||
|
||||
sonar.add_samplers()
|
||||
blep_init()
|
||||
|
||||
NODE_CLASS_MAPPINGS = nodes.NODE_CLASS_MAPPINGS | {
|
||||
"SonarPowerNoise": powernoise.SonarPowerNoiseNode,
|
||||
"SonarPowerFilterNoise": powernoise.SonarPowerFilterNoiseNode,
|
||||
"SonarPowerFilter": powernoise.SonarPowerFilterNode,
|
||||
"SonarPreviewFilter": powernoise.SonarPreviewFilterNode,
|
||||
"FreeUExtremeConfig": freeu_extreme.FreeUExtremeConfigNode,
|
||||
"FreeUExtreme": freeu_extreme.FreeUExtremeNode,
|
||||
}
|
||||
NODE_CLASS_MAPPINGS = nodes.NODE_CLASS_MAPPINGS
|
||||
NODE_DISPLAY_NAME_MAPPINGS = nodes.NODE_DISPLAY_NAME_MAPPINGS
|
||||
NODE_DISPLAY_NAME_MAPPINGS = getattr(nodes, "NODE_DISPLAY_NAME_MAPPINGS", {})
|
||||
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
|
||||
@@ -2,6 +2,78 @@
|
||||
|
||||
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
|
||||
|
||||
## 20250808
|
||||
|
||||
Aside from the `sigmoid` quantile mode change and `SonarShuffledNoise`, these changes should not break workflows. Let me know if you experience anything unusual.
|
||||
|
||||
* The `sigmoid` quantile norm mode was renamed to `sigmoid_keepsign` since that's what it was doing. There is a replacement `sigmoid` quantile norm mode that doesn't care about sign.
|
||||
* Quantile normalization can now take a negative quantile to consider values closest to zero the "extremes". Note: Experimental feature that may not be implemented correctly/subject to change.
|
||||
* `SonarShuffledNoise` node reworked. Unfortunately, this will break workflows. If anyone has a burning need for the old version, let me know and I can bring it back as a separate node. The new approach should be better in general though.
|
||||
* Fixed momentum sampler init parameter passing.
|
||||
* Added more Voronoi noise octave modes.
|
||||
* Added replace_2pt/3pt (and variants) quantile norm result modes that use multiple replacement values.
|
||||
|
||||
## 20250805
|
||||
|
||||
Once again, large set of changes/internal reorganization which may break stuff. If you run into problems or experience anything weird, please create an issue.
|
||||
|
||||
* Added a `SonarResizedNoiseAdv` node that allows more control (and is more useful for models like ACE-Steps where you might want to deal with absolute sizes).
|
||||
* Added a `SonarWaveletCFG` node which allows you use different CFG values for different frequencies.
|
||||
* Added a `SonarCustomNoiseParameters` node that lets you set some parameters as well as override seed/device/dtype.
|
||||
* Added `replace`, `replace_keepsign` and `replace_avoidsign` quantile norm modes.
|
||||
* `SonarBlendedNoise` now has a `custom_noise_mask` input. When connected, it will generate noise with that, put it on a 0-1 scale and use that to control the blend.
|
||||
* Added a `SonarAdvancedVoronoiNoise` node.
|
||||
|
||||
## 20250705
|
||||
|
||||
This is a large set of changes. Please let me know anything doesn't seem to be working properly.
|
||||
|
||||
* Reorganized the node structure. This is an internal change and shouldn't affect users but please let me know if you notice anything weird.
|
||||
* Added a `SonarLatentOperationAdvanced` node which allows more control over when individual latent operations are active and their effects get blended.
|
||||
* Added a `SonarSplitNoiseChain` node. Can be useful if you want to have an item in the chain be a blended.
|
||||
* Added a `SonarLatentOperationNoise` node that can be used to inject noise. You can also use the guided noise node to turn a reference into "noise".
|
||||
* Expanded the functionality of the `SonarWaveletFilteredNoise` node. You can now attach two custom noise inputs to use for the high/low frequency parts of the wavelet as well as blend the wavelets.
|
||||
* Added a `SonarNormalizeNoiseToScale` node that lets you normalize noise to specific value ranges.
|
||||
* Added a `SonarPerDimNoise` node that lets you do stuff like call a noise sampler once per batch item (can be useful for 3D Perlin noise).
|
||||
* Fixed an issue where the normalization parameter wasn't respected. This may change seeds.
|
||||
* Added a `SonarLatentOperationFilteredNoise` node that allows you to run noise through a `LATENT_OPERATION`.
|
||||
* Added a `SonarLatentOperationSetSeed` node that can be used to set the seed (mainly useful for running latent operations that add noise outside of sampling).
|
||||
* Added a `SonarScatternetFilteredNoise` node that uses a scatternet to filter noise. Similar to wavelet filtering. Note: Very experimental, way not work properly.
|
||||
* Fixed an issue with pattern break noise, this may change seeds for workflows using that noise type.
|
||||
|
||||
## 20250627
|
||||
|
||||
* Added `SonarRippleFilteredNoise` node.
|
||||
* Added `SonarApplyLatentOperationCFG` node, similar to the built-in `ApplyLatentOperationCFG` node with scheduling and a lot of different application modes.
|
||||
* Added a `SonarLatentOperationQuantileFilter` node that can be used to apply the quantile normalization functioen to the latent during sampling.
|
||||
* A bunch more quantile normalization modes.
|
||||
* Fixed broken quantile normalization dimension handling. Unfortunately this will likely change seeds.
|
||||
|
||||
## 20250612
|
||||
|
||||
* Reimplemented Collatz noise with many new features. Unfortunately this breaks existing workflows. If anyone misses the old version, let me know and I can add it back in (might do that anyway).
|
||||
* Added actual wavelet noise based on https://en.wikipedia.org/wiki/Wavelet_noise .
|
||||
* Added `reverse_zero`, `scale_down`, `tanh`, `tanh_outliers`, `sigmoid` and `sigmoid_outliers` quantile normalization limit modes.
|
||||
|
||||
## 20250602
|
||||
|
||||
* Fixed broken calculation for Collatz noise.
|
||||
* Added `SonarPatternBreakNoise` node that allows breaking patterns in the noise.
|
||||
* Added `SonarShuffledNoise` node that allows shuffling elements along user-specified dimensions.
|
||||
* Added a strategy option to the `SonarQuantileFilteredNoise` node.
|
||||
* Added variants to Collatz noise. Variant one is maybe similar to the original iteration.
|
||||
* Added `SonarNoiseImage` node that allows generating noisy images or adding noise to existing images.
|
||||
|
||||
## 20250528
|
||||
|
||||
* Added `override_sigma`, `override_sigma_next`, `override_sigma_min` and `override_sigma_max` options that can be set in the `SonarCustomNoiseAdv` node YAML options. This enables using noise generators that require a sigma in stuff like initial noise (for example, Brownian). You will need to manually find and set the correct values yourself.
|
||||
* Added Collatz noise based on the Collatz conjecture. Very experimental, very slow, likely to change and quite possibly just plain bad. But you can try it.
|
||||
|
||||
## 20250505
|
||||
|
||||
* Added `SonarQuantileFilteredNoise` node.
|
||||
* Better compatibility with older Python versions.
|
||||
|
||||
## 20250227
|
||||
|
||||
* Add 5D latent (video models) support for most custom noise types.
|
||||
|
||||
@@ -156,6 +156,12 @@ yh_scales: null
|
||||
|
||||
***
|
||||
|
||||
## `SonarQuantileFilteredNoise`
|
||||
|
||||
Allows quantile normalizing of arbitrary noise generators, works like the `SonarAdvancedDistroNoise` node (see below).
|
||||
|
||||
***
|
||||
|
||||
### `SonarAdvancedDistroNoise`
|
||||
|
||||
See: https://pytorch.org/docs/stable/distributions.html
|
||||
@@ -442,3 +448,59 @@ Only provided if [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh) is ava
|
||||
```
|
||||
|
||||
to flip the sign on the noise and then roll dimension -2 (height) by 50%.
|
||||
|
||||
### `SonarAdvancedVoronoiNoise`
|
||||
|
||||
This node can create multi-octave 3D Voronoi noise (also known as Worley noise). See: https://en.wikipedia.org/wiki/Worley_noise
|
||||
|
||||
Similar to Pyramid and other weird noise types, this noise generally will require mixing with something more normal. The default settings actually just about work with SDXL.
|
||||
|
||||
The node has many options for calculating the distance between the feature points and for processing the output. The modes are entered as a string, you can hover over the widget to get a brief list of possible modes. Both distance and result modes support some common features:
|
||||
|
||||
* You can enter a comma-separated list of modes. This allows using a different mode per octave. If there are more octaves than you have modes defined, the mode will wrap. In other words, if you're generating three octaves and you define two modes then the third octave will use the first mode you defined.
|
||||
* You can enter a `+` (plus symbol) separated list of modes. The modes will be calculated and the result will be the average. Distance modes all have the common parameter `dscale` which defaults to 1 and can be overridden. Result modes use `rscale`. See below for a description on passing parameters.
|
||||
* It's possible to pass parameters to distance and result modes. Example with a result mode: `diff:idx1=0:idx2=1:rscale=0.5`
|
||||
|
||||
Some modes act as wrappers to other modes. All modes will just ignore parameters they don't understand. Each time a submode is called, one level of "_" at the beginning of parameter names is stripped off. It's not user-friendly but this does allow passing parameters to submodes. Since unknown parameters are ignored, you only need to bother with this if the mode that's calling the submode will use that parameter. Dumb example: `gradient_magnitude:name1=diff:name2=gradient_magnitude:_name1=f4:_name2=f4`. `gradient_magnitude` takes two submodes that it calls (specified with `name1` and `name2`). The top-level `gradient_magnitude` will consume the `name1` and `name2` parameters, strip one level of underscores off the parameter names and call the submodes.
|
||||
|
||||
#### Distance Modes
|
||||
|
||||
Modes listed with the defaults for parameters they support. These modes also support `dscale` which defaults to 1 and can be used to manually adjust the scale of the mode result.
|
||||
|
||||
* `euclidean` - Default mode, uses Euclidean distances.
|
||||
* `manhatten` - Uses Manhatten distances.
|
||||
* `chebyshev`
|
||||
* `minkowsi:p=3.0`
|
||||
* `quadratic`
|
||||
* `angle:idx=2` - idxs here range from 0 to 2.
|
||||
* `angle_tanh:idx=2` - Same as `angle` but scales the result with tanh.
|
||||
* `angle_sigmoid:idx=2` - Same as `angle` but scales the result with the sigmoid function.
|
||||
* `fuzz:name=euclidean:fuzz=0.25` - Acts as a wrapper for another mode (specified with `name`). Will perturb the result by `fuzz` percent of the absolute maximum value. Or more simply, randomizes values by +/- `fuzz` percentage so if you set `fuzz=1` you will essentially get pure noise.
|
||||
* `fractal_norm:scale=0.1:multiplier=10.0:mode=sin:name=euclidean` - This acts as a wrapper for another mode. `mode` may be one of `sin`, `cos`. It will adjust the input to the mode it wraps by `scale * sin(input * multiplier)` (assuming `mode=sin`).
|
||||
* `weight:h=1.0:w=1.0:z=0.25:name=euclidean` - This acts as a wrapper for another mode and allows you to scale height/width/z (depth) before calling it.
|
||||
|
||||
#### Result Modes
|
||||
|
||||
Modes listed with the defaults for parameters they support. These modes also support `rscale` which defaults to 1 and can be used to manually adjust the scale of the mode result.
|
||||
|
||||
* `f1` - distance to the closest cell.
|
||||
* `f2`, `f3`, `f4` - Same as `f1` but for the second closest, third closest, etc. Goes up to `f4`.
|
||||
* `f:idx=0` - Allows specifying `f` modes over `f4`. Note that `idx` is zero-based so 0 corresponds to `f1`.
|
||||
* `inv_f1` (through `inv_f4`). For `inv_f1`, the result is `1 / f1` (with a tiny value added to avoid divide by zero).
|
||||
* `inv_f:idx=0` - Like the `f` mode using the formula described above.
|
||||
* `diff:idx1=0:idx2=1` - Zero based indexes where 0 corresponds to `f1`, etc. For the default (`f1` and `f2`) the result is `f2 - f1`.
|
||||
* `diff2:idx1=0:idx2=1` - Similar to `diff` described above, however the result is divided by the two `f` results (plus a tiny addition to avoid divide by zero). For example with the defaults this works out to `(f2 - f1) / (f2 + f1)`.
|
||||
* `cellid` - Returns a discrete value for the area of each cell (diffusion models hate this). You will need to dilute the Voronoi noise a lot to actually use this. It could also possibly be used for masking.
|
||||
* `median_distance`
|
||||
* `fuzz:name=f1:fuzz=0.25` - Works the same as `fuzz` in distance modes. See the description there.
|
||||
* `fractal_norm:scale=0.1:multiplier=10.0:mode=sin:name=diff` - This acts as a wrapper for another mode. `mode` may be one of `sin`, `cos`. It will adjust the input to the mode it wraps by `scale * sin(input * multiplier)` (assuming `mode=sin`).
|
||||
* `ridge:name=diff:exp=-1.0` - Wraps another mode and may enhance cell borders (doesn't seem super useful).
|
||||
* `softmin:temperature=50.0` - Passes the result through softmin and can be used to smooth the output from distance modes like `angle` that may experience abrupt changes as you move through `z`. It also supports a `use_sorted` parameter that will apply this adjustment to the sorted values as well if it's present and set to anything. Higher temperatures will result in less of a smoothing effect.
|
||||
* `gradient_magnitude:name1=f4:name2=f4:padding_mode=replicate` - Wraps two other modes. Seems pretty nice for adding detail when using low numbers of feature points. `padding_mode` can be set to modes that PyTorch's `pad` function supports.
|
||||
|
||||
#### Depth
|
||||
|
||||
|
||||
The node has several `z`-related parameters. `z` here refers to the depth dimension and will (currently) only apply if you're using the same noise sampler more than once. So generally not for initial noise, unless you're doing something unusual.
|
||||
|
||||
When `z_max` is set to 0 the feature points will be reset each time the noise sampler is called. Otherwise it will track the current `z` (depth) and increment it by whatever you specify each time the noise sampler is called. `SonarPerDimNoise` can be useful here if you want depth over dimensions like batch or channels.
|
||||
|
||||
@@ -18,6 +18,8 @@ noise of that type. However you can either schedule the noise type to kick in at
|
||||
* `onef_pinkishgreenish` (50/50 mix of `onef_pinkish` and `onef_greenish`.)
|
||||
* `velvet`
|
||||
* `violet`
|
||||
* `voronoi_mix` - A mix of Voronoi (60%) and Gaussian noise types.
|
||||
* `voronoi_fuzz` - Voronoi noise with distance mode `fuzz:name=angle_tanh:fuzz=0.1`.
|
||||
* `white`
|
||||
|
||||
## Brownian
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
# Wavelet CFG
|
||||
|
||||
A CFG function that lets you use different CFG values for different frequencies.
|
||||
|
||||
Node: `SonarWaveletCFG`
|
||||
|
||||
## Requirements
|
||||
|
||||
You will need to have the `pytorch_wavelets` package installed in your Python environment to use this.
|
||||
|
||||
Link: https://github.com/fbcotter/pytorch_wavelets
|
||||
|
||||
## YAML crash course
|
||||
|
||||
You can skip past this if you already know YAML. Since wavelet CFG definitions are defined with YAML rules,
|
||||
I am putting this section near the beginning.
|
||||
|
||||
First, JSON is valid YAML, so if you know JSON you can use that if you prefer. Since JSON is valid YAML, this also means the structure of YAML documents is the same as a JSON document.
|
||||
|
||||
YAML looks like this:
|
||||
|
||||
```yaml
|
||||
# Comments start with the hash symbol.
|
||||
# YAML will guess the type if you don't do stuff like quote strings, so the item below
|
||||
# will be "value".
|
||||
key1_name: value
|
||||
# The value of a key can also be a set of keys (usually called an object)
|
||||
# YAML uses indentation to control block grouping.
|
||||
key2_name:
|
||||
# A comment
|
||||
subkey1_name: 123
|
||||
# A list of three items, two integers and a string.
|
||||
some_list: [1, 2, "hello"]
|
||||
# You can also specify lists like this:
|
||||
some_other_list:
|
||||
- 1
|
||||
- 2
|
||||
- hello
|
||||
# It's legal to specify the same key multiple times. This just overwrites
|
||||
# whatever the previous value was.
|
||||
subkey1_name: 345
|
||||
```
|
||||
|
||||
YAML item types:
|
||||
|
||||
* String: `"hello there"` or `hello there` (you may need to quote when there are special characters).
|
||||
* Integer: `123`
|
||||
* Floating point value: `1.23`
|
||||
* Object, may be specified in-line like JSON: `{ key: value, key: value }`. Be careful to separate the value from the colon after the key or it may be interpreted incorrectly. It may be a good idea to quote string values if you're using this syntax (and it's necessary if they have special characters or spaces).
|
||||
* List, may be specified in-line like JSON: `[1, 2, "hello there"]`
|
||||
* Null: `null`. If you want the string "null" then you'd need to quote it.
|
||||
* Boolean: `true` and `false`.
|
||||
|
||||
#### Advanced YAML
|
||||
|
||||
YAML also has a number of advanced features like references. You can use this to avoid repeating the same information multiple times. For example:
|
||||
|
||||
```yaml
|
||||
single_value: &single_ref_name [1, 2, 3]
|
||||
# This is the same as other_value: [1, 2, 3]
|
||||
other_value: *single_ref_name
|
||||
|
||||
# You can also do this with objects.
|
||||
reference_block: &ref_name
|
||||
key: value
|
||||
other_key: other_value
|
||||
whatever:
|
||||
# This sets the "whatever" object to be the same as "reference_block"
|
||||
# Note that you can't just do "whatever: *ref_name" to get that effect here.
|
||||
<<: *ref_name
|
||||
# And you can just overwrite the keys you want to change:
|
||||
key: 123
|
||||
# At this point, "whatever" is { key: 123, other_key: other_value }
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
Unfortunately, this isn't very user-friendly and needs to be configured with YAML. This is a relatively basic description of
|
||||
usage. To see all possible options, look at the default configuration definition in the node.
|
||||
|
||||
General information on wavelets from the library I'm using to do wavelet transforms: https://pytorch-wavelets.readthedocs.io/en/latest/index.html
|
||||
|
||||
Trimmed down, the default config looks like this:
|
||||
|
||||
```yaml
|
||||
# This block is used to set the CFG scales.
|
||||
diff:
|
||||
# Scale for the low-frequency band.
|
||||
yl_scale: 5.0
|
||||
|
||||
# Scale for the high-frequency bands.
|
||||
yh_scales: 3.0
|
||||
|
||||
# Sets the wavelet type. DB4 is a good general-purpose wavelet to use.
|
||||
wave: db4
|
||||
|
||||
# Sets the wavelet level.
|
||||
level: 5
|
||||
|
||||
# Set to true if you want to get detailed information dumped to your console.
|
||||
verbose: false
|
||||
```
|
||||
|
||||
Wavelets decompose the value into a low frequency value and a set of high-frequency bands. The number of high-frequency
|
||||
will be equal to the wavelet level, so in this example you will have one low-frequency band and five high-frequency
|
||||
bands to work with. The high-frequency bands are further decomposed into three parts which can also be targeted
|
||||
individually: horizontal, vertical, diagonal.
|
||||
|
||||
The `diff` (or `difference`, whichever you prefer) block is where you set the CFG scales. It's called `difference`
|
||||
because CFG is defined as `uncond + (cond - uncond) * cfg_scale` (`cond` is the positive prompt, `uncond` is
|
||||
negative). So CFG is just the difference between `cond` and `uncond`, multiplied by the CFG scale.
|
||||
|
||||
The example configuration here is using CFG 5 for the low frequency band and CFG 3 for the high frequency bands. It is
|
||||
possible to get even more specific that that. Since we're using `level: 5` here, that means there are five frequency
|
||||
bands that can all be set individually. The high-frequency bands are ordered from fine to coarse detail levels. Example:
|
||||
|
||||
```yaml
|
||||
diff:
|
||||
yl_scale: 5.0
|
||||
# Can also be written: yh_scales: [5.0, 3.0, fill]
|
||||
yh_scales:
|
||||
# Highest/finest band.
|
||||
- 5.0
|
||||
# Decreasing order of detail/frequency.
|
||||
- 3.0
|
||||
- fill
|
||||
```
|
||||
|
||||
The special value `fill` will just repeat the value before it to fill the rest of the bands. Note that if
|
||||
you don't specify the bands or fill then the bands you don't set will use `1.0`. For example with five
|
||||
bands, `[5.0, 3.0, fill]` is the same as `[5.0, 3.0, 3.0, 3.0, 3.0]` while `[5.0, 3.0]` is the same
|
||||
as `[5.0, 3.0, 1.0, 1.0, 1.0]`. You can only use one `fill` per `yh_scales` definition.
|
||||
|
||||
As mentioned, it's also possible to target horizontal, vertical and diagonal bands. You can do this
|
||||
by using a list instead of numeric value for a band definition. **Note**: You need to specify all
|
||||
three bands, `fill` isn't valid here. Example:
|
||||
|
||||
```yaml
|
||||
diff:
|
||||
yl_scale: 5.0
|
||||
# Can also be written: yh_scales: [5.0, 3.0, fill]
|
||||
yh_scales:
|
||||
- [3.0, 3.0, 5.0]
|
||||
- 3.0
|
||||
- fill
|
||||
```
|
||||
|
||||
This example uses CFG 3.0 for horizontal and vertical in finest high-frequency band and CFG 5.0 for
|
||||
diagonal. The remaining high-frequency bands use CFG 3.0.
|
||||
|
||||
## Scheduling CFG
|
||||
|
||||
It's also possible to transition from one set of CFG scales to another over time. The wavelet scales
|
||||
block has an alternative definition format:
|
||||
|
||||
```yaml
|
||||
diff:
|
||||
# One of linear, logarithmic, exponential, half_cosine, sine
|
||||
# Sine mode will hit the peak scales_after values in the middle of the range.
|
||||
schedule: linear
|
||||
|
||||
# One of: sampling, enabled_sampling, sigmas, enabled_sigmas, step, enabled_steps
|
||||
schedule_mode: enabled_sampling
|
||||
|
||||
# When enabled, flips the schedule percentage. This happens before the schedule is applied
|
||||
# or any offset/multiplier stuff. If you want to flip the final result you can do something like
|
||||
# schedule_offset_after: -1.0 and schedule_multiplier_after: -1.0
|
||||
reverse_schedule: false
|
||||
|
||||
scales_start:
|
||||
yl_scale: 5.0
|
||||
yh_scales: 3.0
|
||||
scales_end:
|
||||
yl_scale: 2.0
|
||||
yh_scales: 5.0
|
||||
```
|
||||
|
||||
The way interpolating scales works is we determine a value between 0.0 and 1.0 based on `schedule` and
|
||||
`schedule_mode` and then do linear interpolation (LERP) between the values in `scales_start` and `scales_end`.
|
||||
LERP is just `value_1 * (1.0 - ratio) + value_2 * ratio` so when `ratio` is 1 you get 100% `value_2`,
|
||||
when it's 0.5 you get half of each and when it's 0 you get 100% of `value_1`.
|
||||
|
||||
**Note**: If you have both `scales_start` and toplevel `yl_scale`/`yh_scales` definitions, the
|
||||
scales in `scales_start` will take precedence.
|
||||
|
||||
#### `schedule`
|
||||
|
||||
Current schedule types: `linear`, `logarithmic`, `exponential`, `half_cosine`, `sine`
|
||||
|
||||
Linear just changes by the same amount over time, with the exception
|
||||
of `sine`, the other possible values are similar except the change forms a curve (where it may be slow at first and then
|
||||
accelerate or vice versa). Experiment with them to see what you prefer. `sine` has a somewhat different effect, the
|
||||
percentage of `scales_end` will increase and peak in the middle of the range, then decrease.
|
||||
|
||||
#### `schedule_mode`
|
||||
|
||||
Current modes: `sampling`, `enabled_sampling`, `sigmas`, `enabled_sigmas`, `step`, `enabled_steps`
|
||||
|
||||
* `sampling`: Sampling is a percentage that starts at 0.0 and ends at 1.0 (assuming you're doing txt2img). This isn't really related
|
||||
to the schedule or steps.
|
||||
* `sigmas`: This calculates the difference between the starting sigma and ending sigma as a percentage.
|
||||
* `step`: Not very well tested and may not work (especially with multi-step samplers). The percentage in this case is the percentage
|
||||
steps
|
||||
|
||||
The `_enabled` variants calculate the percentage based on the range that wavelet CFG is enabled for (in other words,
|
||||
the range betwmeen `start_sigma` and `end_sigma`). I'd suggest not using them as it's a lot easier to predict what values
|
||||
will be used when the schedule start/end points aren't also changing.
|
||||
|
||||
***
|
||||
|
||||
In addition to the values described here, there are number of other advanced configuration options that can be used
|
||||
to add/subtract and offset to the calculated percentage value, multiply it, etc. These advanced parameters can be
|
||||
used to do stuff like speed up the transition between config values, keep them within a certain range with minimum/maximum
|
||||
thresholds, etc. See the default YAML config definition in the node to see what is possible.
|
||||
|
||||
|
||||
## Scheduling rules
|
||||
|
||||
Rules may be scheduled using this syntax:
|
||||
|
||||
```yaml
|
||||
rules:
|
||||
- start_sigma: -1.0
|
||||
end_sigma: 5.0
|
||||
diff:
|
||||
yl_scale: 5.0
|
||||
yh_scales: 3.0
|
||||
- start_sigma: -1.0
|
||||
end_sigma: 0.0
|
||||
diff:
|
||||
yl_scale: 2.0
|
||||
yh_scales: 5.0
|
||||
```
|
||||
|
||||
Values from the top-level are valid within a rule. The definitions from the node are added as the first rule,
|
||||
so if you want to only configure stuff in a `rules` block you can just set the start sigma in the node to `0.0` (
|
||||
which will never match). Rules are checked in order and the first matching one is used.
|
||||
|
||||
An alternative method of scheduling rules is to just chain multiple `SonarWaveletCFG` nodes. If you set the fallback
|
||||
mode to `existing` it is also possible to blend the current result with the next matching one (or normal CFG as the
|
||||
case may be).
|
||||
+15
-7
@@ -4,9 +4,10 @@ import contextlib
|
||||
import importlib
|
||||
import sys
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Callable, NamedTuple
|
||||
from typing import TYPE_CHECKING, Any, NamedTuple
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
from types import ModuleType
|
||||
|
||||
|
||||
@@ -32,8 +33,11 @@ class Integrations:
|
||||
return self.modules.get(key)
|
||||
|
||||
@staticmethod
|
||||
def get_custom_node(name: str) -> ModuleType | None:
|
||||
module_key = f"custom_nodes.{name}"
|
||||
def get_custom_node(module_name: str, key: str) -> ModuleType | None:
|
||||
bi_module = sys.modules.get("_blepping_integrations", {}).get(key)
|
||||
if bi_module is not None:
|
||||
return bi_module
|
||||
module_key = f"custom_nodes.{module_name}"
|
||||
with contextlib.suppress(StopIteration):
|
||||
spec = importlib.util.find_spec(module_key)
|
||||
if spec is None:
|
||||
@@ -67,7 +71,7 @@ class Integrations:
|
||||
return
|
||||
self.initialized = True
|
||||
for ih in self.handlers:
|
||||
module = self.get_custom_node(ih.module_name)
|
||||
module = self.get_custom_node(ih.module_name, ih.key)
|
||||
if module is None:
|
||||
continue
|
||||
if ih.handler is not None:
|
||||
@@ -80,7 +84,7 @@ class Integrations:
|
||||
|
||||
|
||||
class SonarIntegrations(Integrations):
|
||||
def __init__(self, *args: list, **kwargs: dict):
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.register_integration("bleh", "ComfyUI-bleh", self.bleh_integration)
|
||||
self.register_integration(
|
||||
@@ -111,13 +115,17 @@ MODULES = SonarIntegrations()
|
||||
|
||||
class IntegratedNode(type):
|
||||
@staticmethod
|
||||
def wrap_INPUT_TYPES(orig_method: Callable, *args: list, **kwargs: dict) -> dict:
|
||||
def wrap_INPUT_TYPES(orig_method: Callable, *args: Any, **kwargs: Any) -> dict:
|
||||
MODULES.initialize()
|
||||
return orig_method(*args, **kwargs)
|
||||
|
||||
def __new__(cls: type, name: str, bases: tuple, attrs: dict) -> object:
|
||||
obj = type.__new__(cls, name, bases, attrs)
|
||||
if hasattr(obj, "INPUT_TYPES"):
|
||||
if hasattr(obj, "INPUT_TYPES") and not getattr(
|
||||
obj.INPUT_TYPES,
|
||||
"_NO_REPLACE",
|
||||
False,
|
||||
):
|
||||
obj.INPUT_TYPES = partial(cls.wrap_INPUT_TYPES, obj.INPUT_TYPES)
|
||||
return obj
|
||||
|
||||
|
||||
@@ -0,0 +1,215 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import random
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import torch
|
||||
|
||||
from . import utils
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Sequence
|
||||
|
||||
|
||||
class SonarLatentOperation:
|
||||
EXTENDED_LATENT_OPERATION = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
start_sigma: float = math.inf,
|
||||
end_sigma: float = 0.0,
|
||||
op=None,
|
||||
):
|
||||
self.start_sigma = start_sigma if start_sigma >= 0 else math.inf
|
||||
self.end_sigma = end_sigma
|
||||
self.op = op
|
||||
|
||||
def enabled(self, sigma: torch.Tensor | float | None = None) -> bool:
|
||||
if isinstance(sigma, torch.Tensor):
|
||||
sigma = sigma.detach().max().cpu().item()
|
||||
return sigma is None or self.end_sigma <= sigma <= self.start_sigma
|
||||
|
||||
def call_op(
|
||||
self,
|
||||
t: torch.Tensor,
|
||||
*args: Any,
|
||||
op=None,
|
||||
**kwargs: Any,
|
||||
) -> torch.Tensor:
|
||||
if op is None:
|
||||
op = self.op
|
||||
if op is None:
|
||||
return t
|
||||
if not getattr(op, "EXTENDED_LATENT_OPERATION", False):
|
||||
return op(latent=t)
|
||||
return op(*args, latent=t, **kwargs)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
latent: torch.Tensor,
|
||||
*,
|
||||
sigma: torch.Tensor | float | None = None,
|
||||
**kwargs: Any,
|
||||
) -> torch.Tensor:
|
||||
if not self.enabled(sigma=sigma):
|
||||
return latent
|
||||
return self.call_op(latent, sigma=sigma, **kwargs)
|
||||
|
||||
|
||||
class SonarLatentOperationAdvanced(SonarLatentOperation):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
blend_mode: str,
|
||||
blend_strength: float,
|
||||
blend_strategy: str,
|
||||
input_multiplier: float,
|
||||
output_multiplier: float,
|
||||
difference_multiplier: float,
|
||||
ops: Sequence,
|
||||
op_alt=None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
super().__init__(**kwargs)
|
||||
self.blend_function = utils.BLENDING_MODES[blend_mode]
|
||||
self.blend_strength = blend_strength
|
||||
self.blend_strategy = blend_strategy
|
||||
self.input_multiplier = input_multiplier
|
||||
self.output_multiplier = output_multiplier
|
||||
self.difference_multiplier = difference_multiplier
|
||||
self.op_alt = op_alt
|
||||
self.ops = ops
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
latent: torch.Tensor,
|
||||
*,
|
||||
sigma: torch.Tensor | float | None = None,
|
||||
**kwargs: Any,
|
||||
) -> torch.Tensor:
|
||||
t = latent
|
||||
enabled = self.enabled(sigma)
|
||||
if not enabled:
|
||||
return (
|
||||
t
|
||||
if self.op_alt is None
|
||||
else self.call_op(t, sigma=sigma, op=self.op_alt, **kwargs)
|
||||
)
|
||||
output = t * self.input_multiplier if self.input_multiplier != 1.0 else t
|
||||
for op in self.ops:
|
||||
output = self.call_op(output, sigma=sigma, op=op, **kwargs)
|
||||
diff = (
|
||||
output * self.output_multiplier if self.output_multiplier == 1.0 else output
|
||||
) - t
|
||||
if self.difference_multiplier != 1.0:
|
||||
diff *= self.difference_multiplier
|
||||
if self.blend_strategy == "difference":
|
||||
return self.blend_function(t, diff, self.blend_strength)
|
||||
if self.blend_strategy == "result":
|
||||
return self.blend_function(t, t + diff, self.blend_strength)
|
||||
raise ValueError(f"Unknown blend strategy: {self.blend_strategy}")
|
||||
|
||||
|
||||
class SonarLatentOperationNoise(SonarLatentOperation):
|
||||
def __init__(
|
||||
self,
|
||||
*args: Any,
|
||||
custom_noise,
|
||||
scale_to_sigma: bool = False,
|
||||
cpu_noise: bool = False,
|
||||
normalize: bool = True,
|
||||
lazy_noise_sampler: bool = False,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.custom_noise = custom_noise
|
||||
self.normalize = normalize
|
||||
self.scale_to_sigma = scale_to_sigma
|
||||
self.cpu_noise = cpu_noise
|
||||
self.lazy_noise_sampler = lazy_noise_sampler
|
||||
self.noise_sampler = None
|
||||
self.cache_id = None
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
latent: torch.Tensor,
|
||||
*,
|
||||
sigma: torch.Tensor | float | None = None,
|
||||
**kwargs: Any,
|
||||
) -> torch.Tensor:
|
||||
t = latent
|
||||
enabled = self.enabled(sigma)
|
||||
if not enabled:
|
||||
return t
|
||||
if isinstance(sigma, float):
|
||||
sigma = t.new_full((1,), sigma)
|
||||
make_ns = not self.lazy_noise_sampler or self.noise_sampler is None
|
||||
sigma_min = sigma_max = sigma_next = None
|
||||
sample_sigmas = (
|
||||
kwargs.get("raw_args", {})
|
||||
.get("model_options", {})
|
||||
.get("transformer_options", {})
|
||||
.get("sample_sigmas")
|
||||
)
|
||||
if sample_sigmas is not None and sigma is not None:
|
||||
guessed_step = (sample_sigmas - sigma).abs().argmin().detach().item()
|
||||
guessed_sigma = sample_sigmas[guessed_step].max().detach().item()
|
||||
if guessed_sigma == sigma and guessed_step + 1 < len(sample_sigmas):
|
||||
sigma_next = sample_sigmas[guessed_step + 1]
|
||||
if self.lazy_noise_sampler and not make_ns:
|
||||
cache_id = (
|
||||
id(sample_sigmas) if isinstance(sample_sigmas, torch.Tensor) else None
|
||||
)
|
||||
make_ns = cache_id is None or cache_id != self.cache_id
|
||||
self.cache_id = cache_id
|
||||
if make_ns and sample_sigmas is not None:
|
||||
sigmas_min = sample_sigmas[sample_sigmas > 0]
|
||||
sigma_min = (
|
||||
sigmas_min.min().detach().item() if torch.any(sigmas_min) else 0.0
|
||||
)
|
||||
del sigmas_min
|
||||
sigma_max = sample_sigmas.max().detach().item()
|
||||
ns = (
|
||||
self.custom_noise.make_noise_sampler(
|
||||
t,
|
||||
sigma_min=sigma_min,
|
||||
sigma_max=sigma_max,
|
||||
normalized=self.normalize,
|
||||
seed=torch.randint(1, 1 << 31, (), device="cpu").item(),
|
||||
cpu=self.cpu_noise,
|
||||
)
|
||||
if make_ns
|
||||
else self.noise_sampler
|
||||
)
|
||||
if make_ns and self.lazy_noise_sampler:
|
||||
self.noise_sampler = ns
|
||||
noise = ns(sigma, sigma if sigma_next is None else sigma_next)
|
||||
if self.scale_to_sigma and sigma is not None:
|
||||
noise *= sigma
|
||||
noise += t
|
||||
return noise
|
||||
|
||||
|
||||
class SonarLatentOperationSetSeed(SonarLatentOperation):
|
||||
def __init__(self, *args: Any, seed: int, restore_rng_state: bool, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.seed = seed
|
||||
self.restore_rng_state = restore_rng_state
|
||||
|
||||
def __call__(self, *args: Any, **kwargs: Any) -> torch.Tensor:
|
||||
if self.restore_rng_state:
|
||||
pyrandst = random.getstate()
|
||||
torchrandst = torch.random.get_rng_state()
|
||||
else:
|
||||
pyrandst = torchrandst = None
|
||||
try:
|
||||
torch.manual_seed(self.seed)
|
||||
random.seed(self.seed)
|
||||
result = super().__call__(*args, **kwargs)
|
||||
finally:
|
||||
if self.restore_rng_state:
|
||||
torch.random.set_rng_state(torchrandst)
|
||||
random.setstate(pyrandst)
|
||||
return result
|
||||
-2562
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,30 @@
|
||||
from . import (
|
||||
base,
|
||||
freeu_extreme,
|
||||
integrations,
|
||||
latent_operations,
|
||||
misc,
|
||||
momentum_samplers,
|
||||
noise_filters,
|
||||
noise_types,
|
||||
powernoise,
|
||||
)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SonarCustomNoise": base.SonarCustomNoiseNode,
|
||||
"SonarCustomNoiseAdv": base.SonarCustomNoiseAdvNode,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
for nm in (
|
||||
freeu_extreme,
|
||||
integrations,
|
||||
latent_operations,
|
||||
misc,
|
||||
momentum_samplers,
|
||||
noise_filters,
|
||||
noise_types,
|
||||
powernoise,
|
||||
):
|
||||
NODE_CLASS_MAPPINGS |= getattr(nm, "NODE_CLASS_MAPPINGS", {})
|
||||
NODE_DISPLAY_NAME_MAPPINGS |= getattr(nm, "NODE_DISPLAY_NAME_MAPPINGS", {})
|
||||
@@ -0,0 +1,294 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import abc
|
||||
from typing import Any
|
||||
|
||||
from .. import noise, utils
|
||||
from ..external import MODULES, IntegratedNode
|
||||
from .base_inputtypes import InputCollection, InputTypes, LazyInputTypes
|
||||
|
||||
try:
|
||||
from comfy_execution import validation as comfy_validation
|
||||
|
||||
if not hasattr(comfy_validation, "validate_node_input"):
|
||||
raise NotImplementedError # noqa: TRY301
|
||||
HAVE_COMFY_UNION_TYPE = comfy_validation.validate_node_input("B", "A,B")
|
||||
except (ImportError, NotImplementedError):
|
||||
HAVE_COMFY_UNION_TYPE = False
|
||||
except Exception as exc: # noqa: BLE001
|
||||
HAVE_COMFY_UNION_TYPE = False
|
||||
print(
|
||||
f"** ComfyUI-sonar: Warning, caught unexpected exception trying to detect ComfyUI union type support. Disabling. Exception: {exc}",
|
||||
)
|
||||
|
||||
NOISE_INPUT_TYPES = frozenset(("SONAR_CUSTOM_NOISE", "OCS_NOISE"))
|
||||
|
||||
if not HAVE_COMFY_UNION_TYPE:
|
||||
|
||||
class Wildcard(str):
|
||||
__slots__ = ("whitelist",)
|
||||
|
||||
@classmethod
|
||||
def __new__(cls, s, *args: Any, whitelist=None, **kwargs: Any):
|
||||
result = super().__new__(s, *args, **kwargs)
|
||||
result.whitelist = whitelist
|
||||
return result
|
||||
|
||||
def __ne__(self, other):
|
||||
return False if self.whitelist is None else other not in self.whitelist
|
||||
|
||||
WILDCARD_NOISE = Wildcard("*", whitelist=NOISE_INPUT_TYPES)
|
||||
else:
|
||||
WILDCARD_NOISE = ",".join(NOISE_INPUT_TYPES)
|
||||
|
||||
|
||||
NOISE_INPUT_TYPES_HINT = (
|
||||
f"The following input types are supported: {', '.join(NOISE_INPUT_TYPES)}"
|
||||
)
|
||||
|
||||
|
||||
class SonarInputCollection(InputCollection):
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._DELEGATE_KEYS = self._DELEGATE_KEYS | frozenset(
|
||||
(
|
||||
"customnoise",
|
||||
"floatpct",
|
||||
"normalizetristate",
|
||||
"selectblend",
|
||||
"selectnoise",
|
||||
"selectscalemode",
|
||||
"yaml",
|
||||
),
|
||||
)
|
||||
|
||||
def yaml(
|
||||
self,
|
||||
name: str = "yaml_parameters",
|
||||
*,
|
||||
tooltip="Allows specifying custom parameters via YAML. Note: When specifying paramaters this way, there is generally not much error checking.",
|
||||
placeholder="# YAML or JSON here",
|
||||
dynamicPrompts=False, # noqa: N803
|
||||
multiline=True,
|
||||
**kwargs: Any,
|
||||
):
|
||||
return self.field(
|
||||
name,
|
||||
"STRING",
|
||||
tooltip=tooltip,
|
||||
placeholder=placeholder,
|
||||
dynamicPrompts=dynamicPrompts,
|
||||
multiline=multiline,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def selectblend(
|
||||
self,
|
||||
name: str = "blend_mode",
|
||||
*,
|
||||
default="lerp",
|
||||
insert_modes=(),
|
||||
tooltip="Mode used for blending. If you have ComfyUI-bleh then you will have access to many more blend modes.",
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
if not MODULES.initialized:
|
||||
raise RuntimeError(
|
||||
"Attempt to get blending modes before integrations were initialized",
|
||||
)
|
||||
return self.field(
|
||||
name,
|
||||
(*insert_modes, *utils.BLENDING_MODES.keys()),
|
||||
default=default,
|
||||
tooltip=tooltip,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def selectscalemode(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
default="nearest-exact",
|
||||
insert_modes=(),
|
||||
tooltip="Mode used for scaling. If you have ComfyUI-bleh then you will have access to many more scale modes.",
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
if not MODULES.initialized:
|
||||
raise RuntimeError(
|
||||
"Attempt to get scale modes before integrations were initialized",
|
||||
)
|
||||
return self.field(
|
||||
name,
|
||||
(*insert_modes, *utils.UPSCALE_METHODS),
|
||||
default=default,
|
||||
tooltip=tooltip,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def selectnoise(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
default="gaussian",
|
||||
insert_types=(),
|
||||
tooltip="Sets the type of noise.",
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
return self.field(
|
||||
name,
|
||||
(*insert_types, *noise.NoiseType.get_names()),
|
||||
default=default,
|
||||
tooltip=tooltip,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def customnoise(
|
||||
self,
|
||||
name: str,
|
||||
add_hint: bool = True, # noqa: FBT001
|
||||
tooltip="Allows connecting a custom noise chain.",
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
if add_hint:
|
||||
tooltip = f"{tooltip}\n{NOISE_INPUT_TYPES_HINT}"
|
||||
return self.field(name, WILDCARD_NOISE, tooltip=tooltip, **kwargs)
|
||||
|
||||
def normalizetristate(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
default="default",
|
||||
tooltip="Controls whether noise is normalized to 1.0 strength.",
|
||||
**kwargs: Any,
|
||||
):
|
||||
return self.field(
|
||||
name,
|
||||
("default", "forced", "disabled"),
|
||||
default=default,
|
||||
tooltip=tooltip,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def floatpct(self, name: str, *, min=0.0, max=1.0, **kwargs: Any): # noqa: A002
|
||||
return self.float(name=name, min=min, max=max, **kwargs)
|
||||
|
||||
|
||||
class SonarInputTypes(InputTypes):
|
||||
_NO_REPLACE = True
|
||||
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
super().__init__(
|
||||
*args,
|
||||
collection_class=SonarInputCollection,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class SonarLazyInputTypes(LazyInputTypes):
|
||||
_NO_REPLACE = True
|
||||
|
||||
def __init__(self, *args: Any, initializers=(MODULES.initialize,), **kwargs: Any):
|
||||
super().__init__(
|
||||
*args,
|
||||
initializers=initializers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class SonarCustomNoiseNodeBase(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "A custom noise item."
|
||||
RETURN_TYPES = ("SONAR_CUSTOM_NOISE",)
|
||||
OUTPUT_TOOLTIPS = ("A custom noise chain.",)
|
||||
CATEGORY = "advanced/noise"
|
||||
FUNCTION = "go"
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_item_class(self):
|
||||
raise NotImplementedError
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda *, include_rescale=True, include_chain=True: (
|
||||
SonarInputTypes()
|
||||
.req_float_factor(
|
||||
default=1.0,
|
||||
tooltip="Scaling factor for the generated noise of this type.",
|
||||
)
|
||||
.req_float_rescale(
|
||||
_skip=not include_rescale,
|
||||
default=0.0,
|
||||
min=0.0,
|
||||
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. When set to 0, rescaling is disabled.",
|
||||
)
|
||||
.opt_customnoise_sonar_custom_noise_opt(
|
||||
_skip=not include_chain,
|
||||
tooltip="Optional input for more custom noise items.",
|
||||
)
|
||||
),
|
||||
initializers=(),
|
||||
)
|
||||
|
||||
def go(
|
||||
self,
|
||||
factor=1.0,
|
||||
rescale=0.0,
|
||||
sonar_custom_noise_opt=None,
|
||||
**kwargs: Any[str, Any],
|
||||
):
|
||||
nis = (
|
||||
sonar_custom_noise_opt.clone()
|
||||
if sonar_custom_noise_opt
|
||||
else noise.CustomNoiseChain()
|
||||
)
|
||||
if factor != 0:
|
||||
nis.add(self.get_item_class()(factor, **kwargs))
|
||||
return (nis if rescale == 0 else nis.rescaled(rescale),)
|
||||
|
||||
|
||||
class NoiseChainInputTypes(SonarInputTypes):
|
||||
def __init__(self, *, parent=SonarCustomNoiseNodeBase, **kwargs: Any):
|
||||
super().__init__(parent=parent, **kwargs)
|
||||
|
||||
|
||||
class NoiseNoChainInputTypes(SonarInputTypes):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
parent=SonarCustomNoiseNodeBase,
|
||||
parent_args=(),
|
||||
parent_kwargs=None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(
|
||||
parent=parent,
|
||||
parent_args=parent_args,
|
||||
parent_kwargs={"include_chain": False, "include_rescale": False}
|
||||
| (parent_kwargs if parent_kwargs is not None else {}),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class SonarCustomNoiseNode(SonarCustomNoiseNodeBase):
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: NoiseChainInputTypes().req_selectnoise_noise_type(
|
||||
tooltip="Sets the type of noise to generate.",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_item_class(cls):
|
||||
return noise.CustomNoiseItem
|
||||
|
||||
|
||||
class SonarCustomNoiseAdvNode(SonarCustomNoiseNode):
|
||||
DESCRIPTION = "A custom noise item allowing advanced YAML parameter input."
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: NoiseChainInputTypes(parent=SonarCustomNoiseNode).opt_yaml(
|
||||
tooltip="Allows specifying custom parameters via YAML. Note: When specifying paramaters this way, there is generally little to no error checking.",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class SonarNormalizeNoiseNodeMixin:
|
||||
@staticmethod
|
||||
def get_normalize(val: str) -> bool | None:
|
||||
return None if val == "default" else val == "forced"
|
||||
@@ -0,0 +1,272 @@
|
||||
# ruff: noqa: A002
|
||||
from __future__ import annotations
|
||||
|
||||
from copy import deepcopy
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Any, TypeVar
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
bi_int = int
|
||||
bi_bool = bool
|
||||
bi_float = float
|
||||
|
||||
|
||||
class InputCollection:
|
||||
_DELEGATE_KEYS = frozenset(
|
||||
(
|
||||
"bool",
|
||||
"boolean",
|
||||
"clip",
|
||||
"conditioning",
|
||||
"field",
|
||||
"float",
|
||||
"image",
|
||||
"int",
|
||||
"latent",
|
||||
"model",
|
||||
"sampler",
|
||||
"seed",
|
||||
"sigmas",
|
||||
"string",
|
||||
"vae",
|
||||
),
|
||||
)
|
||||
|
||||
def __init__(self, **kwargs: Any):
|
||||
self.fields = kwargs
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
splitkey = key.split("_", 1)
|
||||
if len(splitkey) == 1 or splitkey[0] not in self._DELEGATE_KEYS:
|
||||
errstr = f"Unknown attribute {key} for InputCollection"
|
||||
raise AttributeError(errstr)
|
||||
meth = getattr(self, splitkey[0])
|
||||
return partial(meth, splitkey[1]) if len(splitkey) == 2 else meth
|
||||
|
||||
def to_dict(self):
|
||||
return deepcopy(self.fields)
|
||||
|
||||
def clone(self):
|
||||
return InputCollection(**self.to_dict())
|
||||
|
||||
def __len__(self) -> bi_int:
|
||||
return len(self.fields)
|
||||
|
||||
def __contains__(self, key: str) -> bi_bool:
|
||||
return key in self.fields
|
||||
|
||||
def field(
|
||||
self,
|
||||
name: str,
|
||||
type: str | tuple,
|
||||
*,
|
||||
_skip: bi_bool = False,
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
if not _skip:
|
||||
self.fields[name] = (type,) if not kwargs else (type, kwargs)
|
||||
return self
|
||||
|
||||
def string(
|
||||
self,
|
||||
name: str,
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
return self.field(name, "STRING", **kwargs)
|
||||
|
||||
def float(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
step: bi_float = 0.001,
|
||||
min: bi_float = -10000.0,
|
||||
max: bi_float = 10000.0,
|
||||
round: bi_bool = False,
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
return self.field(
|
||||
name,
|
||||
"FLOAT",
|
||||
step=step,
|
||||
min=min,
|
||||
max=max,
|
||||
round=round,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def int(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
min: bi_float = -10000,
|
||||
max: bi_float = 10000,
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
return self.field(
|
||||
name,
|
||||
"INT",
|
||||
min=min,
|
||||
max=max,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def bool(
|
||||
self,
|
||||
name: str,
|
||||
default: bi_bool = False,
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
return self.field(name, "BOOLEAN", default=default, **kwargs)
|
||||
|
||||
boolean = bool # noqa: A003
|
||||
|
||||
def seed(
|
||||
self,
|
||||
name: str = "seed",
|
||||
*,
|
||||
default: bi_int = 0,
|
||||
min: bi_int = 0,
|
||||
max: bi_int = 0xFFFFFFFFFFFFFFFF,
|
||||
tooltip="Seed to use for generated noise",
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
return self.int(
|
||||
name,
|
||||
default=default,
|
||||
min=min,
|
||||
max=max,
|
||||
tooltip=tooltip,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def image(self, name: str = "image", **kwargs: Any) -> InputCollection:
|
||||
return self.field(name, "IMAGE", **kwargs)
|
||||
|
||||
def latent(self, name: str = "latent", **kwargs: Any) -> InputCollection:
|
||||
return self.field(name, "LATENT", **kwargs)
|
||||
|
||||
def conditioning(
|
||||
self,
|
||||
name: str = "conditioning",
|
||||
**kwargs: Any,
|
||||
) -> InputCollection:
|
||||
return self.field(name, "CONDITIONING", **kwargs)
|
||||
|
||||
def model(self, name: str = "model", **kwargs: Any) -> InputCollection:
|
||||
return self.field(name, "MODEL", **kwargs)
|
||||
|
||||
def sigmas(self, name: str = "sigmas", **kwargs: Any) -> InputCollection:
|
||||
return self.field(name, "SIGMAS", **kwargs)
|
||||
|
||||
def sampler(self, name: str = "sampler", **kwargs: Any) -> InputCollection:
|
||||
return self.field(name, "SAMPLER", **kwargs)
|
||||
|
||||
def clip(self, name: str = "clip", **kwargs: Any) -> InputCollection:
|
||||
return self.field(name, "CLIP", **kwargs)
|
||||
|
||||
def vae(self, name: str = "vae", **kwargs: Any) -> InputCollection:
|
||||
return self.field(name, "VAE", **kwargs)
|
||||
|
||||
|
||||
class InputTypes:
|
||||
C = TypeVar("C", bound=type)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
parent=None,
|
||||
parent_field: str | None = "INPUT_TYPES",
|
||||
parent_args=(),
|
||||
parent_kwargs=None,
|
||||
required: dict | C | None = None,
|
||||
optional: dict | C | None = None,
|
||||
collection_class: C = InputCollection,
|
||||
):
|
||||
if parent is not None and parent_field is not None:
|
||||
parent = getattr(parent, parent_field)
|
||||
if isinstance(parent, LazyInputTypes):
|
||||
parent = parent.get_input_types(
|
||||
*parent_args,
|
||||
**({} if parent_kwargs is None else parent_kwargs),
|
||||
)
|
||||
if isinstance(parent, LazyInputTypes):
|
||||
raise TypeError("Unexpected multi-level LazyInputTypes parent!")
|
||||
if required is None:
|
||||
required = {}
|
||||
elif isinstance(required, collection_class):
|
||||
required = required.to_dict()
|
||||
elif not isinstance(required, dict):
|
||||
raise TypeError("Bad type for 'required' parameter.")
|
||||
if optional is None:
|
||||
optional = {}
|
||||
elif isinstance(optional, collection_class):
|
||||
optional = optional.to_dict()
|
||||
elif not isinstance(optional, dict):
|
||||
raise TypeError("Bad type for 'optional' parameter.")
|
||||
if parent is not None:
|
||||
required = parent.required.to_dict() | required
|
||||
optional = parent.optional.to_dict() | optional
|
||||
self.required = collection_class(**required)
|
||||
self.optional = collection_class(**optional)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.required) + len(self.optional)
|
||||
|
||||
def clone(self) -> InputTypes:
|
||||
return InputTypes(required=self.required, optional=self.optional)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"required": self.required.to_dict(),
|
||||
"optional": self.optional.to_dict(),
|
||||
}
|
||||
|
||||
def __call__(self) -> dict:
|
||||
return self.to_dict()
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
if key.startswith("req_"):
|
||||
meth = getattr(self.required, key[4:])
|
||||
elif key.startswith("opt_"):
|
||||
meth = getattr(self.optional, key[4:])
|
||||
else:
|
||||
errstr = f"Unknown attribute {key} for InputTypes"
|
||||
raise AttributeError(errstr)
|
||||
|
||||
def wrapper(*args: Any, **kwargs: Any):
|
||||
meth(*args, **kwargs)
|
||||
return self
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
class LazyInputTypes:
|
||||
def __init__(self, builder: Callable, initializers=()):
|
||||
self._input_types_params = {}
|
||||
self._input_types = None
|
||||
self.builder = builder
|
||||
self.initializers = initializers
|
||||
|
||||
def get_input_types(self, *args: Any, **kwargs: Any):
|
||||
if args or kwargs:
|
||||
args = tuple(args)
|
||||
cache_key = (args, tuple(kwargs.items()))
|
||||
cached = self._input_types_params.get(cache_key)
|
||||
else:
|
||||
cache_key = None
|
||||
cached = self._input_types
|
||||
if cached:
|
||||
return cached
|
||||
for fun in self.initializers:
|
||||
fun()
|
||||
result = self.builder(*args, **kwargs)
|
||||
if not cache_key:
|
||||
self._input_types = result
|
||||
else:
|
||||
self._input_types_params[cache_key] = result
|
||||
return result
|
||||
|
||||
def __call__(self, *args: Any, **kwargs: Any) -> dict:
|
||||
return self.get_input_types(*args, **kwargs)()
|
||||
@@ -2,8 +2,8 @@ from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from . import utils
|
||||
from .external import IntegratedNode
|
||||
from .. import utils
|
||||
from .base import SonarInputTypes, SonarLazyInputTypes
|
||||
from .powernoise import PowerFilter
|
||||
|
||||
|
||||
@@ -29,156 +29,81 @@ def ffilter(x, pfilter, normalization_factor=1.0, cfg_idx=None, filter_cache=Non
|
||||
return x_filt.to(x.dtype, non_blocking=True)
|
||||
|
||||
|
||||
class FreeUExtremeConfigNode(metaclass=IntegratedNode):
|
||||
class FreeUExtremeConfigNode:
|
||||
DESCRIPTION = "Allows setting configuration for FreeU Extreme."
|
||||
RETURN_TYPES = ("FRUX_CONFIG",)
|
||||
FUNCTION = "go"
|
||||
CATEGORY = "model_patches"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"stage_1": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Controls whether this configuration applies to stage 1.",
|
||||
},
|
||||
),
|
||||
"stage_2": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Controls whether this configuration applies to stage 2.",
|
||||
},
|
||||
),
|
||||
"stage_3": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Controls whether this configuration applies to stage 3.",
|
||||
},
|
||||
),
|
||||
"target": (
|
||||
("backbone", "skip", "both"),
|
||||
{
|
||||
"tooltip": "Controls whether this filter applies to backbone or skip layers (or both).",
|
||||
},
|
||||
),
|
||||
"start": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Start time as percentage of sampling this configuration applies to. Inclusive.",
|
||||
},
|
||||
),
|
||||
"end": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "End time as percentage of sampling this configuration applies to. Inclusive.",
|
||||
},
|
||||
),
|
||||
"slice": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Percentage of the layer the FreeU effect is applied to.",
|
||||
},
|
||||
),
|
||||
"slice_offset": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Offset as a percentage the layer is applied to. For example if slice is 0.25 and slice_offset is 0.25 then the filter will apply to the range 25% through 50%.",
|
||||
},
|
||||
),
|
||||
"filter_norm": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": -10.0,
|
||||
"max": 10.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Normalization factor applied to the filter. 1.0 means 100% normalized.",
|
||||
},
|
||||
),
|
||||
"scale": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1,
|
||||
"min": -100.0,
|
||||
"max": 100.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Strength of the effects applied by this configuration.",
|
||||
},
|
||||
),
|
||||
"blend": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": -10.0,
|
||||
"max": 10.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Blends the filtered result based on the specified strength where 1.0 means 100% filtered.",
|
||||
},
|
||||
),
|
||||
"blend_mode": (
|
||||
tuple(utils.BLENDING_MODES.keys()),
|
||||
{
|
||||
"tooltip": "Mode used when blending. Generally only has an effect when blend is set to values other than 0 or 1",
|
||||
},
|
||||
),
|
||||
"hidden_mean": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "You can think of this as FreeU V2 mode.",
|
||||
},
|
||||
),
|
||||
"final": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "When enabled, other configurations won't be considered if this one matched. Otherwise, multiple configurations/filter effects can be stacked.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"sonar_power_filter_opt": (
|
||||
"SONAR_POWER_FILTER",
|
||||
{
|
||||
"tooltip": "Optionally attach a Power Filter here to set filtering parameters.",
|
||||
},
|
||||
),
|
||||
"frux_config_opt": (
|
||||
"FRUX_CONFIG",
|
||||
{
|
||||
"tooltip": "Optionally attach another configuration node here.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: SonarInputTypes()
|
||||
.req_bool_stage_1(
|
||||
default=True,
|
||||
tooltip="Controls whether this configuration applies to stage 1.",
|
||||
)
|
||||
.req_bool_stage_2(
|
||||
default=False,
|
||||
tooltip="Controls whether this configuration applies to stage 2.",
|
||||
)
|
||||
.req_bool_stage_3(
|
||||
default=False,
|
||||
tooltip="Controls whether this configuration applies to stage 3.",
|
||||
)
|
||||
.req_field_target(
|
||||
("backbone", "skip", "both"),
|
||||
default="backbone",
|
||||
tooltip="Controls whether this filter applies to backbone or skip layers (or both).",
|
||||
)
|
||||
.req_floatpct_start(
|
||||
default=0.0,
|
||||
tooltip="Start time as percentage of sampling this configuration applies to. Inclusive.",
|
||||
)
|
||||
.req_floatpct_end(
|
||||
default=1.0,
|
||||
tooltip="End time as percentage of sampling this configuration applies to. Inclusive.",
|
||||
)
|
||||
.req_floatpct_slice(
|
||||
default=1.0,
|
||||
tooltip="Percentage of the layer the FreeU effect is applied to.",
|
||||
)
|
||||
.req_floatpct_slice_offset(
|
||||
default=0.0,
|
||||
tooltip="Offset as a percentage the layer is applied to. For example if slice is 0.25 and slice_offset is 0.25 then the filter will apply to the range 25% through 50%.",
|
||||
)
|
||||
.req_float_filter_norm(
|
||||
default=0.0,
|
||||
min=-10.0,
|
||||
max=10.0,
|
||||
tooltip="Normalization factor applied to the filter. 1.0 means 100% normalized.",
|
||||
)
|
||||
.req_float_scale(
|
||||
default=1.0,
|
||||
tooltip="Strength of the effects applied by this configuration.",
|
||||
)
|
||||
.req_float_blend(
|
||||
default=1.0,
|
||||
tooltip="Blends the filtered result based on the specified strength where 1.0 means 100% filtered.",
|
||||
)
|
||||
.req_selectblend_blend_mode(
|
||||
tooltip="Mode used when blending. Generally only has an effect when blend is set to values other than 0 or 1",
|
||||
)
|
||||
.req_bool_hidden_mean(
|
||||
default=True,
|
||||
tooltip="You can think of this as FreeU V2 mode.",
|
||||
)
|
||||
.req_bool_final(
|
||||
default=True,
|
||||
tooltip="When enabled, other configurations won't be considered if this one matched. Otherwise, multiple configurations/filter effects can be stacked.",
|
||||
)
|
||||
.opt_field_sonar_power_filter_opt(
|
||||
"SONAR_POWER_FILTER",
|
||||
tooltip="Optionally attach a Power Filter here to set filtering parameters.",
|
||||
)
|
||||
.opt_field_frux_config_opt(
|
||||
"FRUX_CONFIG",
|
||||
tooltip="Optionally attach another configuration node here.",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(cls, **kwargs: dict):
|
||||
@@ -330,51 +255,31 @@ class FreeUExtremeConfig:
|
||||
return f"<FRUXConfig: {meh}>"
|
||||
|
||||
|
||||
class FreeUExtremeNode(metaclass=IntegratedNode):
|
||||
class FreeUExtremeNode:
|
||||
DESCRIPTION = "Main FreeU Extreme node. Allows patching a model with the FreeU (V2) effect with more control."
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "go"
|
||||
CATEGORY = "model_patches"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": (
|
||||
"MODEL",
|
||||
{
|
||||
"tooltip": "Model to patch.",
|
||||
},
|
||||
),
|
||||
"cpu_fft": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Controls whether to perform FFT calculations on the CPU. May be necessary for some GPUs that don't have native support for FFT operations at the cost of performance.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"input_config": (
|
||||
"FRUX_CONFIG",
|
||||
{
|
||||
"tooltip": "Allows specifying configuration for input blocks.",
|
||||
},
|
||||
),
|
||||
"middle_config": (
|
||||
"FRUX_CONFIG",
|
||||
{
|
||||
"tooltip": "Allows specifying configuration for middle blocks.",
|
||||
},
|
||||
),
|
||||
"output_config": (
|
||||
"FRUX_CONFIG",
|
||||
{
|
||||
"tooltip": "Allows specifying configuration for output blocks.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
INPUT_TYPES = (
|
||||
SonarInputTypes()
|
||||
.req_model(tooltip="Model to patch.")
|
||||
.req_bool_cpu_fft(
|
||||
tooltip="Controls whether to perform FFT calculations on the CPU. May be necessary for some GPUs that don't have native support for FFT )operations at the cost of performance.",
|
||||
)
|
||||
.opt_field_input_config(
|
||||
"FRUX_CONFIG",
|
||||
tooltip="Allows specifying configuration for input blocks.",
|
||||
)
|
||||
.opt_field_middle_config(
|
||||
"FRUX_CONFIG",
|
||||
tooltip="Allows specifying configuration for middle blocks.",
|
||||
)
|
||||
.opt_field_output_config(
|
||||
"FRUX_CONFIG",
|
||||
tooltip="Allows specifying configuration for output blocks.",
|
||||
)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
@@ -427,3 +332,9 @@ class FreeUExtremeNode(metaclass=IntegratedNode):
|
||||
if ocfg:
|
||||
m.set_model_output_block_patch(out_patch)
|
||||
return (m,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FreeUExtremeConfig": FreeUExtremeConfigNode,
|
||||
"FreeUExtreme": FreeUExtremeNode,
|
||||
}
|
||||
@@ -0,0 +1,288 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from comfy import samplers
|
||||
|
||||
from .. import external, noise
|
||||
from .base import (
|
||||
NoiseNoChainInputTypes,
|
||||
SonarCustomNoiseNodeBase,
|
||||
SonarInputTypes,
|
||||
SonarLazyInputTypes,
|
||||
SonarNormalizeNoiseNodeMixin,
|
||||
)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
|
||||
|
||||
bleh = None
|
||||
|
||||
|
||||
class SonarBlendFilterNoiseNode(
|
||||
SonarCustomNoiseNodeBase,
|
||||
SonarNormalizeNoiseNodeMixin,
|
||||
):
|
||||
DESCRIPTION = "Custom noise type that allows blending and filtering the output of another noise generator using ComfyUI-bleh."
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: NoiseNoChainInputTypes()
|
||||
.req_customnoise_sonar_custom_noise()
|
||||
.req_selectblend(insert_modes=("simple_add",), default="simple_add")
|
||||
.req_field_ffilter(
|
||||
() if bleh is None else tuple(bleh.py.latent_utils.FILTER_PRESETS.keys()),
|
||||
)
|
||||
.req_string_ffilter_custom(default="")
|
||||
.req_float_ffilter_scale(default=1.0)
|
||||
.req_float_ffilter_strength(default=0.0)
|
||||
.req_int_ffilter_threshold(default=1, min=1, max=32)
|
||||
.req_field_enhance_mode(
|
||||
("none",)
|
||||
if bleh is None
|
||||
else ("none", *bleh.py.latent_utils.ENHANCE_METHODS),
|
||||
default="none",
|
||||
)
|
||||
.req_float_enhance_strength(default=0.0)
|
||||
.req_field_affect(("result", "noise", "both"), default="result")
|
||||
.req_normalizetristate_normalize_result()
|
||||
.req_normalizetristate_normalize_noise(),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_item_class(cls):
|
||||
return noise.BlendFilterNoise
|
||||
|
||||
def go(
|
||||
self,
|
||||
*,
|
||||
factor,
|
||||
sonar_custom_noise,
|
||||
blend_mode,
|
||||
ffilter,
|
||||
ffilter_custom,
|
||||
ffilter_scale,
|
||||
ffilter_strength,
|
||||
ffilter_threshold,
|
||||
enhance_mode,
|
||||
enhance_strength,
|
||||
affect,
|
||||
normalize_result,
|
||||
normalize_noise,
|
||||
):
|
||||
if bleh is None:
|
||||
raise RuntimeError("bleh not available")
|
||||
import ast # noqa: PLC0415
|
||||
|
||||
ffilter_custom = ffilter_custom.strip()
|
||||
normalize_result = (
|
||||
None if normalize_result == "default" else normalize_result == "forced"
|
||||
)
|
||||
normalize_noise = (
|
||||
None if normalize_noise == "default" else normalize_noise == "forced"
|
||||
)
|
||||
if ffilter_custom:
|
||||
ffilter = ast.literal_eval(f"[{ffilter_custom}]")
|
||||
elif ffilter == "none":
|
||||
ffilter = None
|
||||
else:
|
||||
ffilter = bleh.py.latent_utils.FILTER_PRESETS[ffilter]
|
||||
return super().go(
|
||||
factor,
|
||||
noise=sonar_custom_noise.clone(),
|
||||
blend_mode=blend_mode,
|
||||
ffilter=ffilter,
|
||||
ffilter_scale=ffilter_scale,
|
||||
ffilter_strength=ffilter_strength,
|
||||
ffilter_threshold=ffilter_threshold,
|
||||
enhance_mode=enhance_mode,
|
||||
enhance_strength=enhance_strength,
|
||||
affect=affect,
|
||||
normalize_noise=self.get_normalize(normalize_noise),
|
||||
normalize_result=self.get_normalize(normalize_result),
|
||||
)
|
||||
|
||||
|
||||
class SonarBlehOpsNoiseNode(
|
||||
SonarCustomNoiseNodeBase,
|
||||
SonarNormalizeNoiseNodeMixin,
|
||||
):
|
||||
DESCRIPTION = (
|
||||
"Custom noise type that allows manipulating noise with ComfyUI-bleh ops."
|
||||
)
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: NoiseNoChainInputTypes()
|
||||
.req_customnoise_sonar_custom_noise()
|
||||
.req_normalizetristate_normalize()
|
||||
.req_yaml_rules(
|
||||
tooltip="Enter rules in the bleh block ops format here.",
|
||||
placeholder="# YAML ops here",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_item_class(cls):
|
||||
return noise.BlehOpsNoise
|
||||
|
||||
def go(
|
||||
self,
|
||||
*,
|
||||
factor,
|
||||
sonar_custom_noise,
|
||||
rules,
|
||||
normalize,
|
||||
):
|
||||
if bleh is None:
|
||||
raise RuntimeError("bleh not available")
|
||||
return super().go(
|
||||
factor,
|
||||
noise=sonar_custom_noise.clone(),
|
||||
rules=bleh.py.nodes.ops.RuleGroup.from_yaml(rules),
|
||||
normalize=normalize,
|
||||
)
|
||||
|
||||
|
||||
restart = None
|
||||
|
||||
|
||||
def KRestartSamplerCustomNoise_INPUT_TYPES_BUILDER():
|
||||
if restart is not None:
|
||||
get_normal_schedulers = getattr(
|
||||
restart.nodes,
|
||||
"get_supported_normal_schedulers",
|
||||
restart.nodes.get_supported_restart_schedulers,
|
||||
)
|
||||
restart_normal_schedulers = get_normal_schedulers()
|
||||
restart_schedulers = restart.nodes.get_supported_restart_schedulers()
|
||||
restart_default_segments = restart.restart_sampling.DEFAULT_SEGMENTS
|
||||
else:
|
||||
restart_default_segments = ""
|
||||
restart_normal_schedulers = restart_schedulers = ()
|
||||
return (
|
||||
SonarInputTypes()
|
||||
.req_model()
|
||||
.req_field_add_noise(("enable", "disable"), default="enable")
|
||||
.req_seed_noise_seed()
|
||||
.req_int_steps(default=20, min=1)
|
||||
.req_float_cfg(default=8.0, min=0.0)
|
||||
.req_sampler()
|
||||
.req_field_scheduler(restart_normal_schedulers)
|
||||
.req_conditioning_positive()
|
||||
.req_conditioning_negative()
|
||||
.req_latent_latent_image()
|
||||
.req_int_start_at_step(default=0, min=0)
|
||||
.req_int_end_at_step(default=10000, min=0)
|
||||
.req_field_return_with_leftover_noise(
|
||||
("disable", "enable"),
|
||||
default="disable",
|
||||
)
|
||||
.req_string_segments(default=restart_default_segments)
|
||||
.req_field_restart_scheduler(restart_schedulers)
|
||||
.req_bool_chunked_mode(default=True)
|
||||
.opt_customnoise_custom_noise_opt(tooltip="Optional custom noise input.")
|
||||
)
|
||||
|
||||
|
||||
class KRestartSamplerCustomNoise:
|
||||
DESCRIPTION = "Restart sampler variant that allows specifying a custom noise type for noise added by restarts."
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(KRestartSamplerCustomNoise_INPUT_TYPES_BUILDER)
|
||||
|
||||
RETURN_TYPES = ("LATENT", "LATENT")
|
||||
RETURN_NAMES = ("output", "denoised_output")
|
||||
FUNCTION = "go"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
model,
|
||||
add_noise,
|
||||
noise_seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
start_at_step,
|
||||
end_at_step,
|
||||
return_with_leftover_noise,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
chunked_mode=False,
|
||||
custom_noise_opt=None,
|
||||
):
|
||||
if restart is None:
|
||||
raise RuntimeError("Restart not available")
|
||||
return restart.restart_sampling.restart_sampling(
|
||||
model,
|
||||
noise_seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
disable_noise=add_noise == "disable",
|
||||
step_range=(start_at_step, end_at_step),
|
||||
force_full_denoise=return_with_leftover_noise != "enable",
|
||||
output_only=False,
|
||||
chunked_mode=chunked_mode,
|
||||
custom_noise=custom_noise_opt.make_noise_sampler
|
||||
if custom_noise_opt
|
||||
else None,
|
||||
)
|
||||
|
||||
|
||||
class RestartSamplerCustomNoise:
|
||||
DESCRIPTION = "Wrapper used to make another sampler Restart compatible. Allows specifying a custom type for noise added by restarts."
|
||||
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
FUNCTION = "go"
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: SonarInputTypes()
|
||||
.req_sampler()
|
||||
.req_bool_chunked_mode(default=True)
|
||||
.opt_customnoise_custom_noise_opt(tooltip="Optional custom noise input."),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(cls, sampler, chunked_mode, custom_noise_opt=None):
|
||||
if restart is None or not hasattr(restart.restart_sampling, "RestartSampler"):
|
||||
raise RuntimeError("Restart not available")
|
||||
restart_options = {
|
||||
"restart_chunked": chunked_mode,
|
||||
"restart_wrapped_sampler": sampler,
|
||||
"restart_custom_noise": None
|
||||
if custom_noise_opt is None
|
||||
else custom_noise_opt.make_noise_sampler,
|
||||
}
|
||||
restart_sampler = samplers.KSAMPLER(
|
||||
restart.restart_sampling.RestartSampler.sampler_function,
|
||||
extra_options=sampler.extra_options | restart_options,
|
||||
inpaint_options=sampler.inpaint_options,
|
||||
)
|
||||
return (restart_sampler,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS |= {
|
||||
"KRestartSamplerCustomNoise": KRestartSamplerCustomNoise,
|
||||
"RestartSamplerCustomNoise": RestartSamplerCustomNoise,
|
||||
"SonarBlendFilterNoise": SonarBlendFilterNoiseNode,
|
||||
"SonarBlehOpsNoise": SonarBlehOpsNoiseNode,
|
||||
}
|
||||
|
||||
|
||||
def init_integrations(integrations):
|
||||
global restart, bleh # noqa: PLW0603
|
||||
restart = integrations.restart
|
||||
bleh = integrations.bleh
|
||||
|
||||
|
||||
external.MODULES.register_init_handler(init_integrations)
|
||||
@@ -0,0 +1,589 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import math
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from .. import utils
|
||||
from ..external import IntegratedNode
|
||||
from ..latent_ops import (
|
||||
SonarLatentOperation,
|
||||
SonarLatentOperationAdvanced,
|
||||
SonarLatentOperationNoise,
|
||||
SonarLatentOperationSetSeed,
|
||||
)
|
||||
from .base import SonarInputTypes, SonarLazyInputTypes
|
||||
from .noise_filters import SonarQuantileFilteredNoiseNode
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import torch
|
||||
|
||||
|
||||
class SonarApplyLatentOperationCFG(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Allows applying a LATENT_OPERATION during sampling. ComfyUI has a few that are builtin and this node pack also includes: SonarLatentOperationQuantileFilter."
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = "latent/advanced/operations"
|
||||
|
||||
FUNCTION = "go"
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: (
|
||||
SonarInputTypes()
|
||||
.req_model()
|
||||
.req_field_mode(
|
||||
(
|
||||
"cond_sub_uncond",
|
||||
"denoised_sub_uncond",
|
||||
"uncond_sub_cond",
|
||||
"denoised",
|
||||
"cond",
|
||||
"uncond",
|
||||
"model_input",
|
||||
),
|
||||
default="cond_sub_uncond",
|
||||
tooltip="cond_sub_uncond is what ComfyUI's latent operations use. The non-sub_uncond modes likely won't work with pred_flip mode enabled. If you have anything but the denoised options selected, this will use pre-CFG, otherwise it will use post-CFG (unless you are using model_input).",
|
||||
)
|
||||
.req_bool_pred_flip_mode(
|
||||
tooltip="Lets you try to apply the latent operation to the noise prediction rather than the image prediction. Doesn't work properly with the non-sub_uncond modes. No real reason it should be better, just something you can try. Note: The noise prediction gets scaled by the sigma first, in case that's useful information.",
|
||||
)
|
||||
.req_bool_require_uncond(
|
||||
tooltip="When enabled, the operation will be skipped if uncond is unavailable. This will also happen if you choose a mode that requires uncond.",
|
||||
)
|
||||
.req_float_start_sigma(
|
||||
default=-1.0,
|
||||
min=-1.0,
|
||||
tooltip="First sigma the effect becomes active. You can set a negative value here to use whatever the model's maximum sigma is.",
|
||||
)
|
||||
.req_float_end_sigma(
|
||||
default=0.0,
|
||||
min=0.0,
|
||||
tooltip="Last sigma the effect is active.",
|
||||
)
|
||||
.req_selectblend_blend_mode(
|
||||
tooltip="Controls how the output of the latent operation is blended with the original result.",
|
||||
)
|
||||
.req_float_blend_strength(
|
||||
default=0.5,
|
||||
tooltip="Strength of the blend. For a normal blend mode like LERP, 1.0 means use 100% of the output from the latent operation, 0.0 means use none of it and only the original value. Note: Blending is applied to the final result of the operations unless you enable immediate_blend, in other words operation_2 sees a full unblended result from operation_1.",
|
||||
)
|
||||
.req_field_blend_scale_mode(
|
||||
(
|
||||
"none",
|
||||
"reverse_sampling",
|
||||
"sampling",
|
||||
"reverse_enabled_range",
|
||||
"enabled_range",
|
||||
"sampling_sin",
|
||||
"enabled_range_sin",
|
||||
),
|
||||
default="reverse_sampling",
|
||||
tooltip="Can be used to scale the blend strength over time. Basically works like blend_strength * scale_factor (see below)\nnone: Just uses the blend_strength you have set.\nreverse_sampling: The opposite of the model sampling percent, so if you're making a new generation, the beginning of sampling will be 1.0 and the end will be 0.0. The recommended option as applying these operations usually works better toward the beginning of sampling.\nsampling: Same as reverse_sampling, except the beginning will be 0.0 and the end will be 1.0.\nreverse_enabled_range: Flipped percentage of the range between start_sigma and end_sigma.\nenabled_range: Percentage of the range between start_sigma and end_sigma.\nsampling_sin: Uses the sampling percentage with the sine function such that blend_strength will hit the peak value in the middle of the range.\nenabled_range_sin: Similar to sampling_sin except it applies to the percentage of the enabled range.",
|
||||
)
|
||||
.req_float_blend_scale_offset(
|
||||
default=0.0,
|
||||
min=-1.0,
|
||||
max=1.0,
|
||||
tooltip="Only applies when blend_scale_mode is not none. Adds the offset to the calculated percentage and then clamps it to be between blend_scale_min and blend_scale_max.",
|
||||
)
|
||||
.req_float_blend_scale_min(
|
||||
default=0.0,
|
||||
tooltip="Only applies when blend_scale_mode is not none. Minimum value for the blend scale percentage. Many blend modes don't tolerate negative values here.",
|
||||
)
|
||||
.req_float_blend_scale_max(
|
||||
default=1.0,
|
||||
tooltip="Only applies when blend_scale_mode is not none. Maximum value for the blend scale percentage. Many blend modes don't tolerate values over 1.0 here.",
|
||||
)
|
||||
.req_bool_immediate_blend(
|
||||
tooltip="You can enable this to do blending immediately after each latent operation is called. Mainly affects the case where you have multiple latent operations connected.",
|
||||
)
|
||||
.opt_field_operation_1(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
|
||||
)
|
||||
.opt_field_operation_2(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
|
||||
)
|
||||
.opt_field_operation_3(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
|
||||
)
|
||||
.opt_field_operation_4(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
|
||||
)
|
||||
.opt_field_operation_5(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_blend_scaling(
|
||||
*,
|
||||
model_sampling: object,
|
||||
scale_mode: str,
|
||||
sigma: float,
|
||||
sigma_t_max: torch.Tensor,
|
||||
start_sigma: float,
|
||||
end_sigma: float,
|
||||
offset: float,
|
||||
min_pct: float,
|
||||
max_pct: float,
|
||||
) -> float | torch.Tensor:
|
||||
if scale_mode == "none":
|
||||
return 1.0
|
||||
if scale_mode in {"sampling", "sampling_sin", "reverse_sampling"}:
|
||||
rev_sampling_pct = (
|
||||
(model_sampling.timestep(sigma_t_max) / 999).clamp(0, 1).detach().item()
|
||||
)
|
||||
result = (
|
||||
1.0 - rev_sampling_pct if scale_mode == "sampling" else rev_sampling_pct
|
||||
)
|
||||
elif scale_mode in {
|
||||
"enabled_range",
|
||||
"enabled_range_sin",
|
||||
"reverse_enabled_range",
|
||||
}:
|
||||
rev_range_pct = (sigma - end_sigma) / (start_sigma - end_sigma)
|
||||
result = (
|
||||
1.0 - rev_range_pct if scale_mode == "enabled_range" else rev_range_pct
|
||||
)
|
||||
else:
|
||||
raise ValueError("Bad blend_scale_mode")
|
||||
if scale_mode.endswith("_sin"):
|
||||
result = math.sin(result * math.pi)
|
||||
return max(min_pct, min(result + offset, max_pct))
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
model,
|
||||
mode: str,
|
||||
pred_flip_mode: bool,
|
||||
require_uncond: bool,
|
||||
start_sigma: float,
|
||||
end_sigma: float,
|
||||
blend_mode: str,
|
||||
blend_strength: float,
|
||||
blend_scale_mode: str,
|
||||
blend_scale_offset: float,
|
||||
blend_scale_min: float,
|
||||
blend_scale_max: float,
|
||||
immediate_blend: bool,
|
||||
operation_1=None,
|
||||
operation_2=None,
|
||||
operation_3=None,
|
||||
operation_4=None,
|
||||
operation_5=None,
|
||||
) -> tuple:
|
||||
if mode == "model_input":
|
||||
if require_uncond:
|
||||
raise ValueError(
|
||||
"require_uncond does not make sense for the model_input mode.",
|
||||
)
|
||||
if pred_flip_mode:
|
||||
raise ValueError(
|
||||
"pred_flip does not make sense for the model_input mode.",
|
||||
)
|
||||
model = model.clone()
|
||||
operations = tuple(
|
||||
SonarLatentOperation(op=o)
|
||||
for o in (operation_1, operation_2, operation_3, operation_4, operation_5)
|
||||
if o is not None
|
||||
)
|
||||
if not operations:
|
||||
return (model,)
|
||||
ms = model.get_model_object("model_sampling")
|
||||
post_cfg_mode = mode in {"denoised", "denoised_sub_uncond"}
|
||||
blend_function = utils.BLENDING_MODES[blend_mode]
|
||||
sigma_max, sigma_min = (
|
||||
ms.sigma_max.detach().item(),
|
||||
ms.sigma_min.detach().item(),
|
||||
)
|
||||
if start_sigma < 0:
|
||||
start_sigma = sigma_max
|
||||
start_sigma = max(sigma_min, min(sigma_max, start_sigma))
|
||||
end_sigma = max(sigma_min, min(sigma_max, end_sigma))
|
||||
if end_sigma > start_sigma:
|
||||
start_sigma, end_sigma = end_sigma, start_sigma
|
||||
if start_sigma == end_sigma:
|
||||
blend_scale_mode = "none"
|
||||
orig_mode = mode
|
||||
|
||||
def patch(args: dict) -> torch.Tensor:
|
||||
nonlocal mode
|
||||
|
||||
x = args["input"]
|
||||
cond_scale = args.get("cond_scale")
|
||||
sigma_t = args["sigma"]
|
||||
sigma_t_max = sigma_t.max()
|
||||
if sigma_t.numel() > 1:
|
||||
shape_pad = (1,) * (x.ndim - sigma_t.ndim)
|
||||
sigma_t = sigma_t.reshape(sigma_t.shape[0], *shape_pad)
|
||||
sigma = sigma_t_max.detach().item()
|
||||
enabled = end_sigma <= sigma <= start_sigma
|
||||
conds_out = args.get("conds_out", ())
|
||||
uncond = (
|
||||
args.get("uncond_denoised")
|
||||
if post_cfg_mode
|
||||
else (conds_out[1] if len(conds_out) > 1 else None)
|
||||
)
|
||||
if uncond is None and (
|
||||
require_uncond
|
||||
or mode in {"uncond", "uncond_sub_cond", "denoised_sub_uncond"}
|
||||
):
|
||||
enabled = False
|
||||
if not enabled:
|
||||
if mode == "model_input":
|
||||
return x
|
||||
return args["denoised"] if post_cfg_mode else conds_out
|
||||
cond = conds_out[0] if not post_cfg_mode and len(conds_out) else None
|
||||
if uncond is None and mode.endswith("_sub_uncond"):
|
||||
mode = orig_mode.split("_", 1)[0]
|
||||
else:
|
||||
mode = orig_mode
|
||||
if mode == "model_input":
|
||||
t1 = x
|
||||
t2 = None
|
||||
elif mode in {"cond", "cond_sub_uncond"}:
|
||||
t1 = cond
|
||||
t2 = uncond if mode == "cond_sub_uncond" else None
|
||||
elif mode in {"uncond", "uncond_sub_cond"}:
|
||||
t1 = uncond
|
||||
t2 = cond if mode == "uncond_sub_cond" else None
|
||||
else:
|
||||
t1 = args["denoised"]
|
||||
t2 = uncond if mode == "denoised_sub_uncond" else None
|
||||
t1_orig = t1
|
||||
if pred_flip_mode:
|
||||
t1 = (x - t1) / sigma_t
|
||||
if t2 is not None:
|
||||
t2 = (x - t2) / sigma_t
|
||||
curr_blend = blend_strength * cls.get_blend_scaling(
|
||||
scale_mode=blend_scale_mode,
|
||||
offset=blend_scale_offset,
|
||||
min_pct=blend_scale_min,
|
||||
max_pct=blend_scale_max,
|
||||
model_sampling=args["model"].model_sampling,
|
||||
start_sigma=start_sigma,
|
||||
end_sigma=end_sigma,
|
||||
sigma=max(sigma_min, min(sigma, sigma_max)),
|
||||
sigma_t_max=sigma_t_max.clamp(sigma_min, sigma_max),
|
||||
)
|
||||
result = t1 - t2 if t2 is not None else t1.clone()
|
||||
for operation in operations:
|
||||
curr_result = operation(
|
||||
result,
|
||||
sigma=sigma,
|
||||
t2=t2,
|
||||
cond=cond,
|
||||
uncond=uncond,
|
||||
cond_scale=cond_scale,
|
||||
raw_args=args,
|
||||
)
|
||||
result = (
|
||||
blend_function(result, curr_result, curr_blend)
|
||||
if immediate_blend
|
||||
else curr_result
|
||||
)
|
||||
if t2 is not None:
|
||||
result += t2
|
||||
if pred_flip_mode:
|
||||
result = x - sigma_t * result
|
||||
if not immediate_blend:
|
||||
result = blend_function(t1_orig, result, curr_blend)
|
||||
if post_cfg_mode or mode == "model_input":
|
||||
return result
|
||||
conds_out = conds_out.copy()
|
||||
conds_out[0 if mode.startswith("cond") else 1] = result
|
||||
return conds_out
|
||||
|
||||
if post_cfg_mode:
|
||||
model.set_model_sampler_post_cfg_function(patch)
|
||||
elif mode == "model_input":
|
||||
|
||||
def patch_wrapper(apply_model, args: dict) -> torch.Tensor:
|
||||
timestep = args["timestep"]
|
||||
patch_args = args | {"sigma": timestep, "model": model.model}
|
||||
return apply_model(patch(patch_args), timestep, **args["c"])
|
||||
|
||||
model.set_model_unet_function_wrapper(patch_wrapper)
|
||||
else:
|
||||
model.set_model_sampler_pre_cfg_function(patch)
|
||||
return (model,)
|
||||
|
||||
|
||||
class SonarLatentOperationQuantileFilter(SonarQuantileFilteredNoiseNode):
|
||||
DESCRIPTION = "Allows applying a quantile normalization function to the latent during sampling. Can be used with Sonar SonarApplyLatentOperationCFG. The just copies most of the parameters from the other quantile normalization node. When it mentions 'noise' it will affect whatever you're applying the latent operation to (denoised, uncond, etc)."
|
||||
RETURN_TYPES = ("LATENT_OPERATION",)
|
||||
CATEGORY = "latent/advanced/operations"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
result = super().INPUT_TYPES()
|
||||
result.pop("optional", None)
|
||||
reqparams = result["required"]
|
||||
for k in (
|
||||
"custom_noise",
|
||||
"reference_noise_opt",
|
||||
"normalize",
|
||||
"normalize_noise",
|
||||
"factor",
|
||||
):
|
||||
reqparams.pop(k, None)
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
quantile: float,
|
||||
dim: str,
|
||||
flatten: bool,
|
||||
norm_power: float,
|
||||
norm_factor: float,
|
||||
strategy: str,
|
||||
norm_power_in: float,
|
||||
sign_mode: str,
|
||||
abs_quantiles: bool,
|
||||
only_outliers: bool,
|
||||
manual_quantiles: str,
|
||||
zero_mean_scale: str,
|
||||
):
|
||||
# TODO: Support an optional reference LATENT_OPERATION.
|
||||
zms, rms, zrs = cls._parse_mean_scales(zero_mean_scale)
|
||||
nq_lo, nq_hi = cls._parse_manual_quantiles(manual_quantiles)
|
||||
qnorm_filter = functools.partial(
|
||||
utils.quantile_normalize,
|
||||
quantile=quantile,
|
||||
dim=None if dim == "global" else int(dim),
|
||||
flatten=flatten,
|
||||
nq_fac=norm_factor,
|
||||
pow_fac=norm_power,
|
||||
strategy=strategy,
|
||||
pow_fac_in=norm_power_in,
|
||||
sign_mode=sign_mode,
|
||||
abs_quantiles=abs_quantiles,
|
||||
only_outliers=only_outliers,
|
||||
nq_lo=nq_lo,
|
||||
nq_hi=nq_hi,
|
||||
zero_mean_scale=zms,
|
||||
restore_mean_scale=rms,
|
||||
zero_result_mean_scale=zrs,
|
||||
)
|
||||
|
||||
return (SonarLatentOperation(op=lambda latent: qnorm_filter(latent)),) # noqa: PLW0108
|
||||
|
||||
|
||||
class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Allows scheduling and other advanced features for latent operations. If you attach the optional extra LATENT_OPERATIONS, they will be called in sequence _before_ blending or output scaling."
|
||||
RETURN_TYPES = ("LATENT_OPERATION",)
|
||||
CATEGORY = "latent/advanced/operations"
|
||||
|
||||
FUNCTION = "go"
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: (
|
||||
SonarInputTypes()
|
||||
.req_float_start_sigma(
|
||||
default=-1.0,
|
||||
min=-1.0,
|
||||
tooltip="First sigma the effect becomes active. You can set a negative value here to use whatever the model's maximum sigma is.",
|
||||
)
|
||||
.req_float_end_sigma(
|
||||
default=0.0,
|
||||
min=0.0,
|
||||
tooltip="Last sigma the effect is active.",
|
||||
)
|
||||
.req_float_input_multiplier(
|
||||
default=1.0,
|
||||
tooltip="Flat multiplier on the input to the latent operation. The multiplied input is *not* used when calculating the difference, it is only passed to the operation.",
|
||||
)
|
||||
.req_float_output_multiplier(
|
||||
default=1.0,
|
||||
tooltip="Flat multiplier on the output from the latent operation. Occurs before blending or calculating the difference.",
|
||||
)
|
||||
.req_float_difference_multiplier(
|
||||
default=1.0,
|
||||
tooltip="Flat multiplier on the difference or change from the original that the operation performed. Occurs after output_multiplier and before blending applies.",
|
||||
)
|
||||
.req_selectblend_blend_mode(
|
||||
default="inject",
|
||||
tooltip="Controls how the change from the operation is combined with the input. The default of inject just adds it scaled by the blend strength. With 1.0 blend strength, this is just using the output from the operation with no change.",
|
||||
)
|
||||
.req_float_blend_strength(
|
||||
default=0.5,
|
||||
tooltip="Strength of the blend.",
|
||||
)
|
||||
.req_field_blend_strategy(
|
||||
("difference", "result"),
|
||||
default="difference",
|
||||
tooltip="Controls whether blending occurs with the difference or changed result after the latent operation.",
|
||||
)
|
||||
.opt_field_operation(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Latent operation to apply.",
|
||||
)
|
||||
.opt_field_operation_alt(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional alternative operation that will be used when the primary one isn't enabled. May be useful in a case when you want one operation between sigma 1.0 and 0.5 and then a difference operation for lower sigmas which is kind of annoying to specify manually (you'd need to do something like configure another operation to start at 0.499999 or something).",
|
||||
)
|
||||
.opt_field_operation_2(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
|
||||
)
|
||||
.opt_field_operation_3(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
|
||||
)
|
||||
.opt_field_operation_4(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
|
||||
)
|
||||
.opt_field_operation_5(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
start_sigma: float,
|
||||
end_sigma: float,
|
||||
input_multiplier: float,
|
||||
output_multiplier: float,
|
||||
difference_multiplier: float,
|
||||
blend_mode: str,
|
||||
blend_strength: float,
|
||||
blend_strategy: str,
|
||||
operation=None,
|
||||
operation_alt=None,
|
||||
operation_2=None,
|
||||
operation_3=None,
|
||||
operation_4=None,
|
||||
operation_5=None,
|
||||
) -> tuple[SonarLatentOperationAdvanced]:
|
||||
operations = tuple(
|
||||
o if isinstance(o, SonarLatentOperation) else SonarLatentOperation(op=o)
|
||||
for o in (operation, operation_2, operation_3, operation_4, operation_5)
|
||||
if o is not None
|
||||
)
|
||||
if operation_alt is not None and not isinstance(
|
||||
operation_alt,
|
||||
SonarLatentOperation,
|
||||
):
|
||||
operation_alt = SonarLatentOperation(op=operation_alt)
|
||||
return (
|
||||
SonarLatentOperationAdvanced(
|
||||
ops=operations,
|
||||
op_alt=operation_alt,
|
||||
start_sigma=start_sigma,
|
||||
end_sigma=end_sigma,
|
||||
input_multiplier=input_multiplier,
|
||||
output_multiplier=output_multiplier,
|
||||
difference_multiplier=difference_multiplier,
|
||||
blend_mode=blend_mode,
|
||||
blend_strength=blend_strength,
|
||||
blend_strategy=blend_strategy,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class SonarLatentOperationNoiseNode(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Latent operation that allows injecting noise."
|
||||
RETURN_TYPES = ("LATENT_OPERATION",)
|
||||
CATEGORY = "latent/advanced/operations"
|
||||
|
||||
FUNCTION = "go"
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: (
|
||||
SonarInputTypes()
|
||||
.req_customnoise_custom_noise()
|
||||
.req_bool_scale_to_sigma(tooltip="Scales the noise to the current sigma.")
|
||||
.req_bool_cpu_noise(
|
||||
tooltip="Controls whether noise is generated on the CPU or GPU. GPU is usually faster but may change seeds for different models of GPU.",
|
||||
)
|
||||
.req_bool_normalize(
|
||||
default=True,
|
||||
tooltip="Controls whether the generated noise is normalized.",
|
||||
)
|
||||
.req_bool_lazy_noise_sampler(
|
||||
default=True,
|
||||
tooltip="When enabled, the latent operation will attempt to cache the noise sampler between calls and only recreate it when necessary. However, there isn't a 100% reliable way for a latent operation to know when sampling starts/ends so if we get it wrong this will lead to non-deterministic generations. I believe the heuristic I'm using to detect this should be reliable but you can disable it if you notice weird results.",
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
custom_noise,
|
||||
scale_to_sigma: bool,
|
||||
cpu_noise: bool,
|
||||
normalize: bool,
|
||||
lazy_noise_sampler: bool,
|
||||
) -> tuple[SonarLatentOperationNoise]:
|
||||
return (
|
||||
SonarLatentOperationNoise(
|
||||
custom_noise=custom_noise,
|
||||
scale_to_sigma=scale_to_sigma,
|
||||
cpu_noise=cpu_noise,
|
||||
normalize=normalize,
|
||||
lazy_noise_sampler=lazy_noise_sampler,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class SonarLatentOperationSetSeedNode(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Latent operation that allows setting a seed. Can be useful for running latent operations that generate noise outside of a normal sampling context (i.e. operations on the initial latent before sampling)."
|
||||
RETURN_TYPES = ("LATENT_OPERATION",)
|
||||
CATEGORY = "latent/advanced/operations"
|
||||
|
||||
FUNCTION = "go"
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: (
|
||||
SonarInputTypes()
|
||||
.req_field_operation("LATENT_OPERATION")
|
||||
.req_seed(
|
||||
tooltip="Seed to set. Note that this is called _every time_ before the operation.",
|
||||
)
|
||||
.req_bool_restore_rng_state(
|
||||
default=False,
|
||||
tooltip="When enabled, the current RNG state is saved just before calling the operation and restored afterwards. In other words, only the latent operation will see the seed you set. Note: This only handles the PyTorch and Python random module states.",
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
operation,
|
||||
seed: int,
|
||||
restore_rng_state: bool,
|
||||
) -> tuple[SonarLatentOperationSetSeed]:
|
||||
return (
|
||||
SonarLatentOperationSetSeed(
|
||||
op=operation,
|
||||
seed=seed,
|
||||
restore_rng_state=restore_rng_state,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SonarApplyLatentOperationCFG": SonarApplyLatentOperationCFG,
|
||||
"SonarLatentOperationQuantileFilter": SonarLatentOperationQuantileFilter,
|
||||
"SonarLatentOperationAdvanced": SonarLatentOperationAdvancedNode,
|
||||
"SonarLatentOperationNoise": SonarLatentOperationNoiseNode,
|
||||
"SonarLatentOperationSetSeed": SonarLatentOperationSetSeedNode,
|
||||
}
|
||||
@@ -0,0 +1,995 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import inspect
|
||||
import math
|
||||
import random
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import yaml
|
||||
from comfy import model_management, samplers
|
||||
from comfy import utils as comfy_utils
|
||||
from tqdm import tqdm
|
||||
|
||||
from .. import noise, utils
|
||||
from ..external import IntegratedNode
|
||||
from ..noise import NoiseType
|
||||
from ..wavelet_cfg import WaveletCFG, WCFGRules
|
||||
from .base import (
|
||||
NoiseChainInputTypes,
|
||||
SonarCustomNoiseNodeBase,
|
||||
SonarInputTypes,
|
||||
SonarLazyInputTypes,
|
||||
SonarNormalizeNoiseNodeMixin,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
try:
|
||||
from comfy import nested_tensor
|
||||
except (ModuleNotFoundError, ImportError):
|
||||
nested_tensor = None
|
||||
|
||||
|
||||
class NoisyLatentLikeNode(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Allows generating noise (and optionally adding it) based on a reference latent. Note: For img2img workflows, you will generally want to enable add_to_latent as well as connecting the model and sigmas inputs."
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
OUTPUT_TOOLTIPS = ("The noisy latent image.",)
|
||||
CATEGORY = "latent/noise"
|
||||
|
||||
FUNCTION = "go"
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: (
|
||||
SonarInputTypes()
|
||||
.req_selectnoise_noise_type(
|
||||
tooltip="Sets the type of noise to generate. Has no effect when the custom_noise_opt input is connected.",
|
||||
)
|
||||
.req_seed()
|
||||
.req_latent(tooltip="Latent used as a reference for generating noise.")
|
||||
.req_float_multiplier(
|
||||
default=1.0,
|
||||
tooltip="Multiplier for the strength of the generated noise. Performed after mul_by_sigmas_opt.",
|
||||
)
|
||||
.req_bool_add_to_latent(
|
||||
tooltip="Add the generated noise to the reference latent rather than adding it to an empty latent. Generally should be enabled for img2img workflows.",
|
||||
)
|
||||
.req_int_repeat_batch(
|
||||
default=1,
|
||||
min=1,
|
||||
tooltip="Repeats the noise generation the specified number of times. For example, if set to two and your reference latent is also batch two you will get a batch of four as output.",
|
||||
)
|
||||
.req_bool_cpu_noise(
|
||||
default=True,
|
||||
tooltip="Controls whether noise will be generated on GPU or CPU. Only affects noise types that support GPU generation (maybe only Brownian).",
|
||||
)
|
||||
.req_bool_normalize(
|
||||
default=True,
|
||||
tooltip="Controls whether the generated noise is normalized to 1.0 strength before scaling. Generally should be left enabled.",
|
||||
)
|
||||
.opt_customnoise_custom_noise_opt()
|
||||
.opt_sigmas_mul_by_sigmas_opt(
|
||||
tooltip="When connected, will scale the generated noise by the first sigma. Must also connect model_opt to enable.",
|
||||
)
|
||||
.opt_model_model_opt(
|
||||
tooltip="Used when mul_by_sigmas_opt is connected, no effect otherwise.",
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
noise_type: str,
|
||||
seed: int | None,
|
||||
latent: dict,
|
||||
multiplier: float = 1.0,
|
||||
add_to_latent=False,
|
||||
repeat_batch=1,
|
||||
cpu_noise=True,
|
||||
normalize=True,
|
||||
custom_noise_opt: object | None = None,
|
||||
mul_by_sigmas_opt: torch.Tensor | None = None,
|
||||
model_opt: object | None = None,
|
||||
):
|
||||
model, sigmas = model_opt, mul_by_sigmas_opt
|
||||
if sigmas is not None and len(sigmas) > 0:
|
||||
if model is None:
|
||||
raise ValueError(
|
||||
"NoisyLatentLike requires a model when sigmas are connected!",
|
||||
)
|
||||
while hasattr(model, "model"):
|
||||
model = model.model
|
||||
latent_scale_factor = model.latent_format.scale_factor
|
||||
model_sigma_max = float(model.model_sampling.sigma_max)
|
||||
first_sigma = float(sigmas[0])
|
||||
max_denoise = (
|
||||
math.isclose(model_sigma_max, first_sigma, rel_tol=1e-05)
|
||||
or first_sigma > model_sigma_max
|
||||
)
|
||||
multiplier *= (
|
||||
float(
|
||||
torch.sqrt(1.0 + sigmas[0] ** 2.0) if max_denoise else sigmas[0],
|
||||
)
|
||||
/ latent_scale_factor
|
||||
)
|
||||
if sigmas is not None and sigmas.numel() > 1:
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
sigma, sigma_next = sigmas[0], sigmas[1]
|
||||
else:
|
||||
sigma_min, sigma_max, sigma, sigma_next = (None,) * 4
|
||||
latent_samples = latent["samples"]
|
||||
orig_device = latent_samples.device
|
||||
want_device = (
|
||||
torch.device("cpu") if cpu_noise else model_management.get_torch_device()
|
||||
)
|
||||
if latent_samples.device != want_device:
|
||||
latent_samples = latent_samples.detach().clone().to(want_device)
|
||||
if custom_noise_opt is not None:
|
||||
ns = custom_noise_opt.make_noise_sampler(
|
||||
latent_samples,
|
||||
sigma_min=sigma_min,
|
||||
sigma_max=sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu_noise,
|
||||
normalized=False,
|
||||
)
|
||||
else:
|
||||
ns = noise.get_noise_sampler(
|
||||
NoiseType[noise_type.upper()],
|
||||
latent_samples,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu_noise,
|
||||
normalized=False,
|
||||
)
|
||||
randst = torch.random.get_rng_state()
|
||||
try:
|
||||
torch.random.manual_seed(seed)
|
||||
result = torch.cat(
|
||||
tuple(ns(sigma, sigma_next) for _ in range(repeat_batch)),
|
||||
dim=0,
|
||||
)
|
||||
finally:
|
||||
torch.random.set_rng_state(randst)
|
||||
result = utils.scale_noise(result, multiplier, normalized=normalize)
|
||||
if add_to_latent:
|
||||
result += latent_samples.repeat(
|
||||
*(repeat_batch if i == 0 else 1 for i in range(latent_samples.ndim)),
|
||||
).to(result)
|
||||
result = result.to(orig_device)
|
||||
return ({"samples": result},)
|
||||
|
||||
|
||||
class SonarNoiseImageNode(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Allows adding noise to an image or generating images full of noise."
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "image"
|
||||
|
||||
FUNCTION = "go"
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: (
|
||||
SonarInputTypes()
|
||||
.req_selectnoise_noise_type(
|
||||
tooltip="Sets the type of noise to generate. Has no effect when the custom_noise_opt input is connected.",
|
||||
)
|
||||
.req_seed()
|
||||
.req_image(tooltip="Image noise will be added to.")
|
||||
.req_float_noise_min(
|
||||
default=0.0,
|
||||
tooltip="Generated noise will be normalized to have values between noise_min and noise_max. If you set them both to the same value then this disables normalization.",
|
||||
)
|
||||
.req_float_noise_max(
|
||||
default=1.0,
|
||||
tooltip="Generated noise will be normalized to have values between noise_min and noise_max. If you set them both to the same value then this disables normalization.",
|
||||
)
|
||||
.req_float_noise_multiplier(
|
||||
default=0.5,
|
||||
tooltip="Multiplier for the strength of the generated noise. This is performed after noise_min/max scaling.",
|
||||
)
|
||||
.req_field_channel_mode(
|
||||
(
|
||||
"RGB",
|
||||
"RGBA",
|
||||
"R",
|
||||
"G",
|
||||
"B",
|
||||
"A",
|
||||
"RA",
|
||||
"GA",
|
||||
"BA",
|
||||
"RG",
|
||||
"RB",
|
||||
"GB",
|
||||
"RGA",
|
||||
"RBA",
|
||||
"GBA",
|
||||
),
|
||||
default="RGB",
|
||||
tooltip="RGBA will also add noise to the alpha channel as well if it exists. Only used for 3 or 4 channel images, for other numbers of channels (i.e. one channel) then all channels will be targeted.",
|
||||
)
|
||||
.req_selectblend(
|
||||
insert_modes=("simple_add",),
|
||||
default="simple_add",
|
||||
tooltip="Controls how the generated noise is combined with the image. simple_add just adds it and blend_strength is ignored in that case.",
|
||||
)
|
||||
.req_float_blend_strength(
|
||||
default=0.5,
|
||||
tooltip="Multiplier for the strength of the generated noise.",
|
||||
)
|
||||
.req_field_overflow_mode(
|
||||
("clamp", "rescale"),
|
||||
default="clamp",
|
||||
tooltip="When set to clamp, values above/below 0, 1 will be set to those values. When set to rescale, the image values will be rescaled such that the minimum value is 0 and the maximum is 1.",
|
||||
)
|
||||
.req_bool_greyscale_mode(
|
||||
tooltip="When set to clamp, values above/below 0, 1 will be set to those values. When set to rescale, the image values will be rescaled such that the minimum value is 0 and the maximum is 1.",
|
||||
)
|
||||
.req_bool_pure_noise_mode(
|
||||
tooltip="When enabled, the original image is only used for its shape and you will be adding noise to an image full of zeros (black), suitable for creating pure noise images.",
|
||||
)
|
||||
.req_field_dtype(
|
||||
("default", "float32", "float64", "float16", "bfloat16"),
|
||||
default="default",
|
||||
tooltip="When set to default it will use the same type as the input tensor (probably float32). You can manually set the dtype if you want, though it likely isn't going to matter. Using dtypes with limited range (float16, bfloat16) isn't recommended.",
|
||||
)
|
||||
.req_bool_cpu_noise(
|
||||
default=True,
|
||||
tooltip="Controls whether noise will be generated on GPU or CPU.",
|
||||
)
|
||||
.req_bool_normalize(
|
||||
default=True,
|
||||
tooltip="Controls whether the generated noise is normalized to 1.0 strength before scaling. Generally should be left enabled.",
|
||||
)
|
||||
.opt_customnoise_custom_noise_opt(
|
||||
tooltip="Allows connecting a custom noise chain. When connected, noise_type has no effect.",
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
noise_type: str,
|
||||
seed: int,
|
||||
image: torch.Tensor,
|
||||
noise_multiplier: float,
|
||||
noise_min: float,
|
||||
noise_max: float,
|
||||
channel_mode: str,
|
||||
blend_mode: str,
|
||||
blend_strength: float,
|
||||
overflow_mode: str,
|
||||
greyscale_mode: bool,
|
||||
dtype: str,
|
||||
pure_noise_mode: bool,
|
||||
cpu_noise: bool,
|
||||
normalize: bool,
|
||||
custom_noise_opt: object | None = None,
|
||||
):
|
||||
sigma_min, sigma_max, sigma, sigma_next = (None,) * 4
|
||||
orig_image = image = (
|
||||
torch.zeros_like(image) if pure_noise_mode else image.detach().clone()
|
||||
)
|
||||
if image.ndim == 3:
|
||||
image = image.unsqueeze(0)
|
||||
elif image.ndim != 4:
|
||||
errstr = (
|
||||
f"Expected image tensor with 3 or 4 dimensions, got {image.ndim}",
|
||||
)
|
||||
raise ValueError(errstr)
|
||||
blend_function = (
|
||||
utils.BLENDING_MODES[blend_mode]
|
||||
if blend_mode != "simple_add"
|
||||
else lambda a, b, _t: a + b
|
||||
)
|
||||
if noise_min > noise_max:
|
||||
noise_min, noise_max = noise_max, noise_min
|
||||
image = image.movedim(-1, 1)
|
||||
channels = image.shape[1]
|
||||
channel_map = {"R": 0, "B": 1, "G": 2, "A": 3}
|
||||
channel_mode = channel_mode.upper()
|
||||
if channels == 3 or channels == 4: # noqa: PLR1714
|
||||
channel_targets = tuple(
|
||||
channel_map[c]
|
||||
for c in "RGBA"
|
||||
if c in channel_mode and channel_map[c] < channels
|
||||
)
|
||||
else:
|
||||
channel_targets = tuple(range(channels))
|
||||
want_device = (
|
||||
torch.device("cpu") if cpu_noise else model_management.get_torch_device()
|
||||
)
|
||||
image = image.to(
|
||||
device=want_device,
|
||||
dtype={
|
||||
"float32": torch.float32,
|
||||
"float64": torch.float64,
|
||||
"bfloat16": torch.bfloat16,
|
||||
"float16": torch.float16,
|
||||
}.get(
|
||||
dtype,
|
||||
image.dtype,
|
||||
),
|
||||
)
|
||||
pyrandst = random.getstate()
|
||||
randst = torch.random.get_rng_state()
|
||||
try:
|
||||
random.seed(seed)
|
||||
torch.random.manual_seed(seed)
|
||||
if custom_noise_opt is not None:
|
||||
ns = custom_noise_opt.make_noise_sampler(
|
||||
image,
|
||||
sigma_min=sigma_min,
|
||||
sigma_max=sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu_noise,
|
||||
normalized=normalize,
|
||||
)
|
||||
else:
|
||||
ns = noise.get_noise_sampler(
|
||||
NoiseType[noise_type.upper()],
|
||||
image,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu_noise,
|
||||
normalized=normalize,
|
||||
)
|
||||
result = ns(sigma, sigma_next)
|
||||
finally:
|
||||
torch.random.set_rng_state(randst)
|
||||
random.setstate(pyrandst)
|
||||
del ns
|
||||
result = utils.scale_noise(result, normalized=True)
|
||||
if greyscale_mode:
|
||||
result = result.mean(dim=1, keepdim=True).expand(image.shape).contiguous()
|
||||
if noise_max != 0 and noise_min != noise_max: # noqa: PLR1714
|
||||
result = utils.normalize_to_scale(result, noise_min, noise_max)
|
||||
result *= noise_multiplier
|
||||
image[:, channel_targets, ...] = blend_function(
|
||||
image[:, channel_targets, ...],
|
||||
result[:, channel_targets, ...],
|
||||
blend_strength,
|
||||
)
|
||||
if overflow_mode == "rescale":
|
||||
image = utils.normalize_to_scale(image, 0.0, 1.0)
|
||||
else:
|
||||
image = image.clip_(0, 1)
|
||||
image = image.movedim(1, -1).to(
|
||||
device=orig_image.device,
|
||||
dtype=orig_image.dtype,
|
||||
)
|
||||
return (image,)
|
||||
|
||||
|
||||
class CustomNOISE:
|
||||
def __init__(
|
||||
self,
|
||||
custom_noise,
|
||||
seed,
|
||||
*,
|
||||
cpu_noise=True,
|
||||
normalize=True,
|
||||
multiplier=1.0,
|
||||
):
|
||||
self.custom_noise = custom_noise
|
||||
self.seed = seed
|
||||
self.cpu_noise = cpu_noise
|
||||
self.normalize = normalize
|
||||
self.multiplier = multiplier
|
||||
|
||||
def _sample_noise(self, latent_image, seed, sampler_idx: int = 0):
|
||||
if self.multiplier == 0.0:
|
||||
return torch.zeros_like(latent_image)
|
||||
n_samplers = len(self.custom_noise)
|
||||
sampler_idx = sampler_idx % n_samplers
|
||||
result = (
|
||||
self.custom_noise[sampler_idx]
|
||||
.make_noise_sampler(
|
||||
latent_image,
|
||||
None,
|
||||
None,
|
||||
seed=seed,
|
||||
cpu=self.cpu_noise,
|
||||
normalized=self.normalize,
|
||||
)(None, None)
|
||||
.to(
|
||||
device="cpu",
|
||||
dtype=latent_image.dtype,
|
||||
)
|
||||
)
|
||||
if result.layout != latent_image.layout:
|
||||
if latent_image.layout == torch.sparse_coo:
|
||||
return result.to_sparse()
|
||||
errstr = f"Cannot handle latent layout {type(latent_image.layout).__name__}"
|
||||
raise NotImplementedError(errstr)
|
||||
if self.multiplier != 1.0:
|
||||
result *= self.multiplier
|
||||
return result
|
||||
|
||||
def generate_noise(self, input_latent):
|
||||
latent_image = input_latent["samples"]
|
||||
orig_type = type(latent_image)
|
||||
# print(f"\nNEST? {latent_image.is_nested}, have={nested_tensor is not None}")
|
||||
if nested_tensor is not None and latent_image.is_nested:
|
||||
nested = True
|
||||
nested_parts = latent_image.unbind()
|
||||
# print(f"NEST: {tuple(p.shape for p in nested_parts)}")
|
||||
else:
|
||||
nested = False
|
||||
nested_parts = (latent_image,)
|
||||
|
||||
batch_inds = input_latent.get("batch_index")
|
||||
torch.manual_seed(self.seed)
|
||||
random.seed(self.seed)
|
||||
if batch_inds is None:
|
||||
noise_parts = tuple(
|
||||
self._sample_noise(nested_parts[i], self.seed, sampler_idx=i)
|
||||
for i in range(len(nested_parts))
|
||||
)
|
||||
return (
|
||||
orig_type(comfy_utils.pack_latents(noise_parts)[0])
|
||||
if nested
|
||||
else noise_parts[0]
|
||||
)
|
||||
|
||||
batch_size = latent_image.shape[0]
|
||||
unique_inds, inverse_inds = np.unique(batch_inds, return_inverse=True)
|
||||
use_idxs = (idx for idx in range(unique_inds[-1] + 1) if idx in unique_inds)
|
||||
use_idxs = {idx: inverse_inds[uidx] for uidx, idx in enumerate(use_idxs)}
|
||||
n_use_idxs = len(use_idxs)
|
||||
result_parts = tuple(
|
||||
torch.empty(
|
||||
(n_use_idxs, *np.shape[1:]),
|
||||
dtype=latent_image.dtype,
|
||||
device=latent_image.device,
|
||||
)
|
||||
for np in nested_parts
|
||||
)
|
||||
for idx in range(unique_inds[-1] + 1):
|
||||
sample_idx = idx % batch_size
|
||||
for nidx in range(len(nested_parts)):
|
||||
sample = nested_parts[nidx][sample_idx].unsqueeze(0)
|
||||
noise = self._sample_noise(sample, self.seed + idx, sampler_idx=nidx)
|
||||
batch_out_idx = use_idxs.get(idx)
|
||||
if batch_out_idx is not None:
|
||||
result = result_parts[nidx]
|
||||
result[batch_out_idx : batch_out_idx + 1] = noise[:1]
|
||||
return (
|
||||
orig_type(comfy_utils.pack_latents(result_parts)[0])
|
||||
if nested
|
||||
else result_parts[0]
|
||||
)
|
||||
|
||||
|
||||
class SonarToComfyNOISENode(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Allows converting SONAR_CUSTOM_NOISE to NOISE (used by SamplerCustomAdvanced and possibly other custom samplers). The extra alt inputs are used if the latent is a nested tensor and ignored otherwise. Audio/video models like LTX and MiniMax H3 used nested tensors (order video then audio). Connected custom noise inputs will be used in order. NOTE: This node does not work with noise types that depend on sigma (Brownian, ScheduledNoise, etc) unless you manually set a sigma via other nodes."
|
||||
RETURN_TYPES = ("NOISE",)
|
||||
CATEGORY = "sampling/custom_sampling/noise"
|
||||
FUNCTION = "go"
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: (
|
||||
SonarInputTypes()
|
||||
.req_customnoise_custom_noise(
|
||||
tooltip="Custom noise type to convert.",
|
||||
)
|
||||
.req_seed(tooltip="Seed to use for generated noise.")
|
||||
.req_bool_cpu_noise(
|
||||
default=False,
|
||||
tooltip="Controls whether noise is generated on CPU or GPU.",
|
||||
)
|
||||
.req_bool_normalize(
|
||||
default=True,
|
||||
tooltip="Controls whether generated noise is normalized to 1.0 strength.",
|
||||
)
|
||||
.req_float_multiplier(
|
||||
default=1.0,
|
||||
tooltip="Simple multiplier applied to noise after all other scaling and normalization effects. If set to 0, no noise will be generated (same as disabling noise).",
|
||||
)
|
||||
.opt_customnoise_alt_custom_noise_1(
|
||||
tooltip="Optional custom noise. See the node description.",
|
||||
)
|
||||
.opt_customnoise_alt_custom_noise_2(
|
||||
tooltip="Optional custom noise. See the node description.",
|
||||
)
|
||||
.opt_customnoise_alt_custom_noise_3(
|
||||
tooltip="Optional custom noise. See the node description.",
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
custom_noise,
|
||||
seed,
|
||||
cpu_noise=True,
|
||||
normalize=True,
|
||||
multiplier=1.0,
|
||||
alt_custom_noise_1=None,
|
||||
alt_custom_noise_2=None,
|
||||
alt_custom_noise_3=None,
|
||||
):
|
||||
noises = tuple(
|
||||
cn.clone()
|
||||
for cn in (
|
||||
custom_noise,
|
||||
alt_custom_noise_1,
|
||||
alt_custom_noise_2,
|
||||
alt_custom_noise_3,
|
||||
)
|
||||
if cn is not None
|
||||
)
|
||||
return (
|
||||
CustomNOISE(
|
||||
noises,
|
||||
seed,
|
||||
cpu_noise=cpu_noise,
|
||||
normalize=normalize,
|
||||
multiplier=multiplier,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class SamplerNodeConfigOverride(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Allows overriding paramaters for a SAMPLER. Only parameters that particular sampler supports will be applied, so for example setting ETA will have no effect for non-ancestral Euler."
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: (
|
||||
SonarInputTypes()
|
||||
.req_sampler()
|
||||
.req_float_eta(
|
||||
default=1.0,
|
||||
tooltip="Basically controls the ancestralness of the sampler. When set to 0, you will get a non-ancestral (or SDE) sampler.",
|
||||
)
|
||||
.req_float_s_noise(
|
||||
default=1.0,
|
||||
tooltip="Multiplier for noise added during ancestral or SDE sampling.",
|
||||
)
|
||||
.req_float_s_churn(
|
||||
default=0.0,
|
||||
tooltip="Churn was the predececessor of ETA. Only used by a few types of samplers (notably Euler non-ancestral). Not used by any ancestral or SDE samplers.",
|
||||
)
|
||||
.req_float_r(
|
||||
default=0.5,
|
||||
tooltip="Used by dpmpp_sde (and perhaps a few other SDE samplers).",
|
||||
)
|
||||
.req_field_sde_solver(
|
||||
("midpoint", "heun"),
|
||||
tooltip="Solver used by dpmpp_2m_sde.",
|
||||
)
|
||||
.req_bool_cpu_noise(
|
||||
default=True,
|
||||
tooltip="Controls whether noise is generated on CPU or GPU.",
|
||||
)
|
||||
.req_bool_normalize(
|
||||
default=True,
|
||||
tooltip="Controls whether generated noise is normalized to 1.0 strength.",
|
||||
)
|
||||
.opt_selectnoise_noise_type(
|
||||
insert_types=("DEFAULT",),
|
||||
default="DEFAULT",
|
||||
tooltip="Noise type used during ancestral or SDE sampling. DEFAULT will use the default for the attached sampler. Only used when the custom noise input is not connected.",
|
||||
)
|
||||
.opt_customnoise_custom_noise_opt(
|
||||
tooltip="Optional input for custom noise used during ancestral or SDE sampling. When connected, the built-in noise_type selector is ignored.",
|
||||
)
|
||||
.opt_yaml()
|
||||
),
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(
|
||||
self,
|
||||
*,
|
||||
sampler,
|
||||
eta,
|
||||
s_noise,
|
||||
s_churn,
|
||||
r,
|
||||
sde_solver,
|
||||
cpu_noise=True,
|
||||
noise_type=None,
|
||||
custom_noise_opt=None,
|
||||
normalize=True,
|
||||
yaml_parameters="",
|
||||
):
|
||||
sampler_kwargs = {
|
||||
"s_noise": s_noise,
|
||||
"eta": eta,
|
||||
"s_churn": s_churn,
|
||||
"r": r,
|
||||
"solver_type": sde_solver,
|
||||
}
|
||||
if yaml_parameters:
|
||||
extra_params = yaml.safe_load(yaml_parameters)
|
||||
if extra_params is None:
|
||||
pass
|
||||
elif not isinstance(extra_params, dict):
|
||||
raise ValueError(
|
||||
"SamplerConfigOverride: yaml_parameters must either be null or an object",
|
||||
)
|
||||
else:
|
||||
sampler_kwargs |= extra_params
|
||||
sampler_function = functools.update_wrapper(
|
||||
functools.partial(
|
||||
self.sampler_function,
|
||||
override_sampler_cfg={
|
||||
"sampler": sampler,
|
||||
"noise_type": NoiseType[noise_type.upper()]
|
||||
if noise_type not in {None, "DEFAULT"}
|
||||
else None,
|
||||
"custom_noise": custom_noise_opt,
|
||||
"sampler_kwargs": sampler_kwargs,
|
||||
"cpu_noise": cpu_noise,
|
||||
"normalize": normalize,
|
||||
},
|
||||
),
|
||||
sampler.sampler_function,
|
||||
)
|
||||
return (
|
||||
samplers.KSAMPLER(
|
||||
sampler_function,
|
||||
extra_options=sampler.extra_options.copy(),
|
||||
inpaint_options=sampler.inpaint_options.copy(),
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def sampler_function(
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
*args: Any,
|
||||
override_sampler_cfg: dict[str, Any] | None = None,
|
||||
noise_sampler: Callable | None = None,
|
||||
extra_args: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> torch.Tensor:
|
||||
if not override_sampler_cfg:
|
||||
raise ValueError("Override sampler config missing!")
|
||||
if extra_args is None:
|
||||
extra_args = {}
|
||||
cfg = override_sampler_cfg
|
||||
sampler, sampler_kwargs, noise_type, custom_noise, cpu, normalize = (
|
||||
cfg["sampler"],
|
||||
cfg["sampler_kwargs"],
|
||||
cfg.get("noise_type"),
|
||||
cfg.get("custom_noise"),
|
||||
cfg.get("cpu_noise", True),
|
||||
cfg.get("normalize", True),
|
||||
)
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
sig = inspect.signature(sampler.sampler_function)
|
||||
params = sig.parameters
|
||||
if "noise_sampler" in params:
|
||||
seed = extra_args.get("seed")
|
||||
if custom_noise is not None:
|
||||
noise_sampler = custom_noise.make_noise_sampler(
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu,
|
||||
normalized=normalize,
|
||||
)
|
||||
elif noise_type is not None:
|
||||
noise_sampler = noise.get_noise_sampler(
|
||||
noise_type,
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=cpu,
|
||||
normalized=normalize,
|
||||
)
|
||||
kwargs |= {k: v for k, v in sampler_kwargs.items() if k in params}
|
||||
if "noise_sampler" in params:
|
||||
kwargs["noise_sampler"] = noise_sampler
|
||||
return sampler.sampler_function(
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
*args,
|
||||
extra_args=extra_args,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class SonarSplitNoiseChainNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin):
|
||||
DESCRIPTION = "Custom noise type that allows splitting off a new chain. This can be useful if you want a link in the chain to be a blended type."
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: (
|
||||
NoiseChainInputTypes()
|
||||
.req_normalizetristate_normalize(
|
||||
tooltip="Controls whether the generated noise is normalized to 1.0 strength.",
|
||||
)
|
||||
.opt_customnoise_custom_noise()
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_item_class(cls):
|
||||
return noise.BlendedNoise
|
||||
|
||||
def go(
|
||||
self,
|
||||
*,
|
||||
factor,
|
||||
rescale,
|
||||
sonar_custom_noise_opt=None,
|
||||
normalize,
|
||||
custom_noise=None,
|
||||
):
|
||||
return super().go(
|
||||
factor,
|
||||
rescale=rescale,
|
||||
sonar_custom_noise_opt=sonar_custom_noise_opt,
|
||||
blend_function=lambda a, _b, _t: a,
|
||||
normalize=self.get_normalize(normalize),
|
||||
custom_noise_1=custom_noise,
|
||||
custom_noise_2=None,
|
||||
noise_2_percent=0.0,
|
||||
)
|
||||
|
||||
|
||||
class SonarWaveletCFGNode(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Wavelet CFG function that allows you to apply different CFG strength to different frequencies."
|
||||
CATEGORY = "model_patches"
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "go"
|
||||
|
||||
_yaml_placeholder = """# YAML or JSON here.
|
||||
# I recommend reading the documentation at https://github.com/blepping/ComfyUI-sonar/docs/waveletcfg.md
|
||||
# For wavelet information, see: https://pytorch-wavelets.readthedocs.io/en/latest/index.html
|
||||
|
||||
# You may override the fields from the node like start_sigma here.
|
||||
|
||||
# This section is basically the CFG scale. (All scales sections use the same format.)
|
||||
difference:
|
||||
# Scale for the low frequency components.
|
||||
yl_scale: 5.0
|
||||
|
||||
# Scale (or scales) for high frequency components.
|
||||
# This can be scalar or a list or list of lists.
|
||||
# List example:
|
||||
# yh_scales:
|
||||
# - [1, 2, 3]
|
||||
# - fill
|
||||
# - 5
|
||||
# You can separately apply a scale to items equal to the wavelet level. Levels go from fine to coarse.
|
||||
# If the item is a list, the three items correspond to horizontal, vertical, diagonal for DWT. (DTCWT has 6.)
|
||||
# You can have one "fill" item, this will replicate the item before it however many times is necessary to
|
||||
# match the wavelet level.
|
||||
yh_scales: 3.0
|
||||
|
||||
# You can optionally include a scales_end block with yl_scale/yh_scales.
|
||||
# to interpolate from the toplevel scales (can also be in a scales_start blockx if you prefer).
|
||||
|
||||
# scales_end:
|
||||
# yl_scale: 1.0
|
||||
# yh_scales: 1.0
|
||||
|
||||
# The following scheduling parameters only apply if scales_end exists.
|
||||
|
||||
# One of linear, logarithmic, exponential, half_cosine, sine
|
||||
# Sine mode will hit the peak scales_after values in the middle of the range.
|
||||
schedule: linear
|
||||
|
||||
# One of: sampling, enabled_sampling, sigmas, enabled_sigmas, step, enabled_steps
|
||||
schedule_mode: sampling
|
||||
|
||||
# When enabled, flips the schedule percentage. This happens before the schedule is applied
|
||||
# or any offset/multiplier stuff. If you want to flip the final result you can do something like
|
||||
# schedule_offset_after: -1.0 and schedule_multiplier_after: -1.0
|
||||
reverse_schedule: false
|
||||
|
||||
# Added to the percentage before the schedule function is applied.
|
||||
schedule_offset: 0.0
|
||||
|
||||
# Applied to the percentage before the schedule function (but after the offset).
|
||||
schedule_multiplier: 1.0
|
||||
|
||||
# Added to the percentage after the schedule function is applied.
|
||||
schedule_offset_after: 0.0
|
||||
|
||||
# Applied to the percentage after the schedule function (but after the offset).
|
||||
schedule_multiplier_after: 1.0
|
||||
|
||||
# Min/max for the final calculated percent. Must be between 0 and 1.
|
||||
schedule_min: 0.0
|
||||
schedule_max: 1.0
|
||||
|
||||
# If you're a crazy person, you can use non-standard blend modes for interpolating
|
||||
# the scales. Not recommended.
|
||||
blend_mode: lerp
|
||||
|
||||
|
||||
# Wavelet type
|
||||
wave: db4
|
||||
|
||||
# Wavelet level
|
||||
level: 5
|
||||
|
||||
### Start of advanced options
|
||||
|
||||
# Mode used for padding
|
||||
padding_mode: symmetric
|
||||
|
||||
# Mutually exclusive with DTCWT mode.
|
||||
use_1d_dwt: false
|
||||
|
||||
# Enables DTCWT mode.
|
||||
use_dtcwt: false
|
||||
|
||||
# Configuration for DTCWT, only relevant when enabled.
|
||||
biort: near_sym_a
|
||||
qshift: qshift_a
|
||||
|
||||
# It's also possible to set these wavelet options with an "inv_"
|
||||
# prefix: mode, biort, qshift, wave, padding_mode
|
||||
|
||||
# One of: noise_norm, noise, denoised
|
||||
# Normal CFG uses denoised mode. noise_norm divides by the current sigma, noise just uses the raw noise prediction.
|
||||
target_mode: denoised
|
||||
|
||||
# Can be used to scale cond before the difference is calculated.
|
||||
cond:
|
||||
yl_scale: 1.0
|
||||
yh_scales: 1.0
|
||||
|
||||
# Can be used to scale uncond before the difference is calculated.
|
||||
uncond:
|
||||
yl_scale: 1.0
|
||||
yh_scales: 1.0
|
||||
|
||||
# Can be used to scale the final result after blending.
|
||||
final:
|
||||
yl_scale: 1.0
|
||||
yh_scales: 1.0
|
||||
|
||||
# Uses float64 for the wavelets/scaling/blending operations.
|
||||
# It doesn't seem to hurt performance much, but you can disable it if you want.
|
||||
high_precision_mode: true
|
||||
|
||||
# Inject is just addition which is usually what you want. The normal CFG function is:
|
||||
# uncond + (cond - uncond) * cfg_scale
|
||||
difference_blend_mode: inject
|
||||
difference_blend_strength: 1.0
|
||||
|
||||
# Per-rule value, can be enabled to spam your console with information when
|
||||
# rules activate, dump exactly what high/low scales are used, etc.
|
||||
verbose: false
|
||||
|
||||
# You may include a rules block which is a list of these configuration definitions.
|
||||
# Include start_sigma/end_sigma parameters. The first matching definition will be used.
|
||||
# rules:
|
||||
# - start_sigma: -1.0
|
||||
"""
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda _yaml_placeholder=_yaml_placeholder: (
|
||||
SonarInputTypes()
|
||||
.req_model()
|
||||
.req_float_start_sigma(
|
||||
default=-1.0,
|
||||
min=-1.0,
|
||||
tooltip="First sigma wavelet CFG will be used.",
|
||||
)
|
||||
.req_float_end_sigma(
|
||||
default=0.0,
|
||||
min=0.0,
|
||||
tooltip="Last sigma wavelet CFG will be used.",
|
||||
)
|
||||
.req_field_fallback_mode(
|
||||
("existing", "own"),
|
||||
default="existing",
|
||||
tooltip="Existing mode uses whatever CFG function existed set when this model patch was applied. Own mode does the CFG calculation on its own. The scale will be whatever you set in your guider or sampler.",
|
||||
)
|
||||
.req_selectblend_blend_mode(
|
||||
tooltip="Controls how the result from wavelet CFG is blended with normal CFG. The default of LERP with strength 1.0 uses 100% wavelet CFG.",
|
||||
)
|
||||
.req_float_blend_strength(
|
||||
default=1.0,
|
||||
tooltip="Controls how the result from wavelet CFG is blended with normal CFG. The default of LERP with strength 1.0 uses 100% wavelet CFG.",
|
||||
)
|
||||
.req_yaml(default=_yaml_placeholder)
|
||||
.opt_field_operation_cond(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional latent operation that will be applied to cond. Note: Latent operations only apply if a rule matches.",
|
||||
)
|
||||
.opt_field_operation_uncond(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional latent operation that will be applied to uncond. Note: Latent operations only apply if a rule matches.",
|
||||
)
|
||||
.opt_field_operation_fallback_cfg(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional latent operation that will be applied to the fallback (non-wavelet) CFG result. Note: Latent operations only apply if a rule matches.",
|
||||
)
|
||||
.opt_field_operation_wavelet_cfg(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional latent operation that will be applied to wavelet CFG result. Note: Latent operations only apply if a rule matches.",
|
||||
)
|
||||
.opt_field_operation_result(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional latent operation that will be applied to the final result, after wavelet and normal CFG are potentially blended. Note: Latent operations only apply if a rule matches.",
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
model: object,
|
||||
start_sigma: float,
|
||||
end_sigma: float,
|
||||
fallback_mode: str,
|
||||
blend_mode: str,
|
||||
blend_strength: float,
|
||||
yaml_parameters: str,
|
||||
operation_cond: Callable | None = None,
|
||||
operation_uncond: Callable | None = None,
|
||||
operation_fallback_cfg: Callable | None = None,
|
||||
operation_wavelet_cfg: Callable | None = None,
|
||||
operation_result: Callable | None = None,
|
||||
_override_rules_dict: dict | None = None,
|
||||
) -> tuple[object]:
|
||||
if start_sigma < 0:
|
||||
start_sigma = math.inf
|
||||
if _override_rules_dict is not None:
|
||||
wavelet_params = _override_rules_dict.copy()
|
||||
else:
|
||||
wavelet_params = yaml.safe_load(yaml_parameters)
|
||||
rules = WCFGRules.build(
|
||||
**(
|
||||
{
|
||||
"start_sigma": start_sigma,
|
||||
"end_sigma": end_sigma,
|
||||
"fallback_existing": fallback_mode == "existing",
|
||||
"blend_mode": blend_mode,
|
||||
"blend_strength": blend_strength,
|
||||
}
|
||||
| wavelet_params
|
||||
),
|
||||
)
|
||||
if len(rules) and rules[0].verbose:
|
||||
tqdm.write(f"\nWCFG: Using rules: {rules}\n")
|
||||
model = model.clone()
|
||||
model.set_model_sampler_cfg_function(
|
||||
WaveletCFG(
|
||||
existing_cfg=model.model_options.get("sampler_cfg_function"),
|
||||
rules=rules,
|
||||
operation_cond=operation_cond,
|
||||
operation_uncond=operation_uncond,
|
||||
operation_fallback_cfg=operation_fallback_cfg,
|
||||
operation_wavelet_cfg=operation_wavelet_cfg,
|
||||
operation_result=operation_result,
|
||||
),
|
||||
)
|
||||
return (model,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"NoisyLatentLike": NoisyLatentLikeNode,
|
||||
"SamplerConfigOverride": SamplerNodeConfigOverride,
|
||||
"SONAR_CUSTOM_NOISE to NOISE": SonarToComfyNOISENode,
|
||||
"SonarNoiseImage": SonarNoiseImageNode,
|
||||
"SonarSplitNoiseChain": SonarSplitNoiseChainNode,
|
||||
"SonarWaveletCFG": SonarWaveletCFGNode,
|
||||
}
|
||||
@@ -0,0 +1,249 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from comfy import samplers
|
||||
|
||||
from ..external import IntegratedNode
|
||||
from ..noise import NoiseType
|
||||
from ..sonar import (
|
||||
GuidanceConfig,
|
||||
GuidanceType,
|
||||
HistoryType,
|
||||
SonarConfig,
|
||||
SonarDPMPPSDE,
|
||||
SonarEuler,
|
||||
SonarEulerAncestral,
|
||||
)
|
||||
from .base import SonarInputTypes, SonarLazyInputTypes
|
||||
|
||||
|
||||
class GuidanceConfigNode(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Allows specifying extended guidance parameters for Sonar samplers."
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: SonarInputTypes()
|
||||
.req_float_factor(
|
||||
default=0.01,
|
||||
min=-2.0,
|
||||
max=2.0,
|
||||
tooltip="Controls the strength of the guidance. You'll generally want to use fairly low values here.",
|
||||
)
|
||||
.req_field_guidance_type(
|
||||
tuple(t.name.lower() for t in GuidanceType),
|
||||
default="linear",
|
||||
tooltip="Method to use when calculating guidance. When set to linear, will simply LERP the guidance at the specified strength. When set to Euler, will do a Euler step toward the guidance instead.",
|
||||
)
|
||||
.req_int_start_step(
|
||||
default=0,
|
||||
min=0,
|
||||
tooltip="First zero-based step the guidance is active.",
|
||||
)
|
||||
.req_int_end_step(
|
||||
default=9999,
|
||||
min=0,
|
||||
tooltip="Last zero-based step the guidance is active.",
|
||||
)
|
||||
.req_latent(tooltip="Latent to use as a reference for guidance."),
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("SONAR_GUIDANCE_CFG",)
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
FUNCTION = "make_guidance_cfg"
|
||||
|
||||
@classmethod
|
||||
def make_guidance_cfg(
|
||||
cls,
|
||||
guidance_type,
|
||||
factor,
|
||||
start_step,
|
||||
end_step,
|
||||
latent,
|
||||
):
|
||||
return (
|
||||
GuidanceConfig(
|
||||
guidance_type=GuidanceType[guidance_type.upper()],
|
||||
factor=factor,
|
||||
start_step=start_step,
|
||||
end_step=end_step,
|
||||
latent=latent.get("samples"),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class SamplerNodeSonarBase:
|
||||
DESCRIPTION = "Sonar - momentum based sampler node."
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: SonarInputTypes()
|
||||
.req_float_momentum(
|
||||
default=0.95,
|
||||
min=-0.5,
|
||||
max=2.5,
|
||||
tooltip="How much of the normal result to keep during sampling. 0.95 means 95% normal, 5% from history. When set to 1.0 effectively disables momentum.",
|
||||
)
|
||||
.req_float_momentum_hist(
|
||||
default=0.75,
|
||||
min=-1.5,
|
||||
max=1.5,
|
||||
tooltip="How much of the existing history to leave at each update. 0.75 means keep 75%, mix in 25% of the new result.",
|
||||
)
|
||||
.req_field_momentum_init(
|
||||
tuple(t.name for t in HistoryType),
|
||||
default="ZERO",
|
||||
tooltip="Initial value used for momentum history. ZERO - history starts zeroed out. RAND - History is initialized with a random value. SAMPLE - History is initialized from the latent at the start of sampling.",
|
||||
)
|
||||
.req_float_direction(
|
||||
default=1.0,
|
||||
min=-30.0,
|
||||
max=15.0,
|
||||
tooltip="Multiplier applied to the result of normal sampling.",
|
||||
)
|
||||
.req_field_rand_init_noise_type(
|
||||
tuple(NoiseType.get_names(skip=(NoiseType.BROWNIAN,))),
|
||||
default="gaussian",
|
||||
tooltip="Noise type to use when momentum_init is set to RANDOM.",
|
||||
)
|
||||
.opt_field_guidance_cfg_opt(
|
||||
"SONAR_GUIDANCE_CFG",
|
||||
tooltip="Optional input for extended guidance parameters.",
|
||||
),
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
|
||||
class SamplerNodeSonarEuler(SamplerNodeSonarBase):
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
@classmethod
|
||||
def get_sampler(
|
||||
cls,
|
||||
*,
|
||||
momentum,
|
||||
momentum_hist,
|
||||
momentum_init,
|
||||
direction,
|
||||
rand_init_noise_type,
|
||||
guidance_cfg_opt=None,
|
||||
):
|
||||
cfg = SonarConfig(
|
||||
momentum=momentum,
|
||||
init=HistoryType[momentum_init.upper()],
|
||||
momentum_hist=momentum_hist,
|
||||
direction=direction,
|
||||
rand_init_noise_type=NoiseType[rand_init_noise_type.upper()],
|
||||
guidance=guidance_cfg_opt,
|
||||
)
|
||||
return (samplers.KSAMPLER(SonarEuler.sampler, {"sonar_config": cfg}),)
|
||||
|
||||
|
||||
class SamplerNodeSonarEulerAncestral(SamplerNodeSonarEuler):
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: SonarInputTypes(parent=SamplerNodeSonarEuler)
|
||||
.req_float_s_noise(
|
||||
default=1.0,
|
||||
tooltip="Multiplier for noise added during ancestral or SDE sampling.",
|
||||
)
|
||||
.req_float_eta(
|
||||
default=1.0,
|
||||
tooltip="Basically controls the ancestralness of the sampler. When set to 0, you will get a non-ancestral (or SDE) sampler.",
|
||||
)
|
||||
.req_selectnoise_noise_type(
|
||||
tooltip="Noise type used during ancestral or SDE sampling. Only used when the custom noise input is not connected.",
|
||||
)
|
||||
.opt_customnoise_custom_noise_opt(
|
||||
tooltip="Optional input for custom noise used during ancestral or SDE sampling. When connected, the built-in noise_type selector is ignored.",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_sampler(
|
||||
cls,
|
||||
*,
|
||||
momentum,
|
||||
momentum_hist,
|
||||
momentum_init,
|
||||
direction,
|
||||
rand_init_noise_type,
|
||||
noise_type,
|
||||
eta,
|
||||
s_noise,
|
||||
guidance_cfg_opt=None,
|
||||
custom_noise_opt=None,
|
||||
):
|
||||
cfg = SonarConfig(
|
||||
momentum=momentum,
|
||||
init=HistoryType[momentum_init.upper()],
|
||||
momentum_hist=momentum_hist,
|
||||
direction=direction,
|
||||
rand_init_noise_type=NoiseType[rand_init_noise_type.upper()],
|
||||
noise_type=NoiseType[noise_type.upper()],
|
||||
custom_noise=custom_noise_opt.clone() if custom_noise_opt else None,
|
||||
guidance=guidance_cfg_opt,
|
||||
)
|
||||
return (
|
||||
samplers.KSAMPLER(
|
||||
SonarEulerAncestral.sampler,
|
||||
{
|
||||
"sonar_config": cfg,
|
||||
"eta": eta,
|
||||
"s_noise": s_noise,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class SamplerNodeSonarDPMPPSDE(SamplerNodeSonarEulerAncestral):
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: SonarInputTypes(
|
||||
parent=SamplerNodeSonarEulerAncestral,
|
||||
).req_selectnoise_noise_type(default="brownian"),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_sampler(
|
||||
cls,
|
||||
*,
|
||||
momentum,
|
||||
momentum_hist,
|
||||
momentum_init,
|
||||
direction,
|
||||
rand_init_noise_type,
|
||||
noise_type,
|
||||
eta,
|
||||
s_noise,
|
||||
guidance_cfg_opt=None,
|
||||
custom_noise_opt=None,
|
||||
):
|
||||
cfg = SonarConfig(
|
||||
momentum=momentum,
|
||||
init=HistoryType[momentum_init.upper()],
|
||||
momentum_hist=momentum_hist,
|
||||
direction=direction,
|
||||
rand_init_noise_type=NoiseType[rand_init_noise_type.upper()],
|
||||
noise_type=NoiseType[noise_type.upper()],
|
||||
custom_noise=custom_noise_opt.clone() if custom_noise_opt else None,
|
||||
guidance=guidance_cfg_opt,
|
||||
)
|
||||
return (
|
||||
samplers.KSAMPLER(
|
||||
SonarDPMPPSDE.sampler,
|
||||
{
|
||||
"sonar_config": cfg,
|
||||
"eta": eta,
|
||||
"s_noise": s_noise,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SamplerSonarEuler": SamplerNodeSonarEuler,
|
||||
"SamplerSonarEulerA": SamplerNodeSonarEulerAncestral,
|
||||
"SamplerSonarDPMPPSDE": SamplerNodeSonarDPMPPSDE,
|
||||
"SonarGuidanceConfig": GuidanceConfigNode,
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -7,6 +7,7 @@ from __future__ import annotations
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from typing import Any
|
||||
|
||||
import comfy
|
||||
import folder_paths
|
||||
@@ -16,16 +17,16 @@ from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler
|
||||
from PIL import Image
|
||||
from torch import Tensor
|
||||
|
||||
from .nodes import (
|
||||
from ..noise import CustomNoiseItemBase
|
||||
from ..utils import scale_noise
|
||||
from .base import (
|
||||
NOISE_INPUT_TYPES_HINT,
|
||||
WILDCARD_NOISE,
|
||||
NoiseChainInputTypes,
|
||||
SonarCustomNoiseNodeBase,
|
||||
SonarInputTypes,
|
||||
SonarNormalizeNoiseNodeMixin,
|
||||
)
|
||||
from .noise import CustomNoiseItemBase
|
||||
from .utils import scale_noise
|
||||
|
||||
# ruff: noqa: ANN003, FBT001, FBT002
|
||||
|
||||
PREVIEW_FORMAT = comfy.latent_formats.SD15()
|
||||
|
||||
@@ -86,7 +87,7 @@ class ChannelMixer:
|
||||
channel_mixer /= channel_mixer.norm(dim=1, keepdim=True)
|
||||
return channel_mixer
|
||||
|
||||
def to(self, *args: list, **kwargs: dict):
|
||||
def to(self, *args: Any, **kwargs: Any):
|
||||
if self.mixer is not None:
|
||||
self.mixer = self.mixer.to(*args, **kwargs)
|
||||
return self
|
||||
@@ -100,7 +101,7 @@ class ChannelMixer:
|
||||
noise = self.mixer @ noise.swapaxes(0, 1).reshape(c, -1)
|
||||
return noise.reshape(c, b, h, w).swapaxes(1, 0)
|
||||
|
||||
def __call__(self, *args: list, **kwargs: dict):
|
||||
def __call__(self, *args: Any, **kwargs: Any):
|
||||
return self.apply(*args, **kwargs)
|
||||
|
||||
|
||||
@@ -295,7 +296,14 @@ class PowerFilter:
|
||||
|
||||
|
||||
class PowerNoiseItem(CustomNoiseItemBase):
|
||||
def __init__(self, factor, *, channel_correlation, power_filter=None, **kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
factor,
|
||||
*,
|
||||
channel_correlation,
|
||||
power_filter=None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
if isinstance(channel_correlation, str):
|
||||
channel_correlation = torch.tensor(
|
||||
tuple(
|
||||
@@ -365,6 +373,7 @@ class PowerNoiseItem(CustomNoiseItemBase):
|
||||
x: Tensor,
|
||||
sigma_min: float | None,
|
||||
sigma_max: float | None,
|
||||
*,
|
||||
seed: int | None,
|
||||
cpu: bool = True,
|
||||
normalized=True,
|
||||
@@ -461,7 +470,15 @@ def rfft2_to_fft2(x):
|
||||
|
||||
|
||||
class PowerFilterNoiseItem(PowerNoiseItem):
|
||||
def __init__(self, factor, *, noise, normalize_noise, normalize_result, **kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
factor,
|
||||
*,
|
||||
noise,
|
||||
normalize_noise,
|
||||
normalize_result,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(
|
||||
factor,
|
||||
noise=noise.clone(),
|
||||
@@ -480,6 +497,7 @@ class PowerFilterNoiseItem(PowerNoiseItem):
|
||||
x: Tensor,
|
||||
sigma_min: float | None,
|
||||
sigma_max: float | None,
|
||||
*,
|
||||
seed: int | None,
|
||||
cpu: bool = True,
|
||||
normalized=True,
|
||||
@@ -540,122 +558,69 @@ class PowerFilterNoiseItem(PowerNoiseItem):
|
||||
class SonarPowerNoiseNode(SonarCustomNoiseNodeBase):
|
||||
DESCRIPTION = "Custom noise type that applies a filter to generated noise."
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls, *args: list, **kwargs: dict):
|
||||
result = super().INPUT_TYPES(*args, **kwargs)
|
||||
result["required"] |= {
|
||||
"time_brownian": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Controls whether brownian noise is used when mix isn't 1.0.",
|
||||
},
|
||||
),
|
||||
"alpha": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": -5.0,
|
||||
"max": 5.0,
|
||||
"step": 0.001,
|
||||
"round": False,
|
||||
"tooltip": "Values above 0 will amplify low frequencies, negative values will amplify high frequencies.",
|
||||
},
|
||||
),
|
||||
"max_freq": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.7071,
|
||||
"min": 0.0,
|
||||
"max": 0.7071,
|
||||
"step": 0.001,
|
||||
"round": False,
|
||||
"tooltip": "Maximum frequency to pass through the filter.",
|
||||
},
|
||||
),
|
||||
"min_freq": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 0.7071,
|
||||
"step": 0.001,
|
||||
"round": False,
|
||||
"tooltip": "Minimum frequency to pass through the filter.",
|
||||
},
|
||||
),
|
||||
"stretch": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.01,
|
||||
"max": 100,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Stretches the filter's shape by the specified factor.",
|
||||
},
|
||||
),
|
||||
"rotate": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": -90,
|
||||
"max": 90,
|
||||
"step": 5,
|
||||
"round": False,
|
||||
"tooltip": "Rotates the filter.",
|
||||
},
|
||||
),
|
||||
"pnorm": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 2,
|
||||
"min": 0.125,
|
||||
"max": 100,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Factor used for cushioning the band-pass region.",
|
||||
},
|
||||
),
|
||||
"mix": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.001,
|
||||
"round": False,
|
||||
"tooltip": "Controls the ratio of filtered noise. For example, 0.75 means 75% noise with the filter effects applied, 25% raw noise.",
|
||||
},
|
||||
),
|
||||
"common_mode": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": -100.0,
|
||||
"max": 100.0,
|
||||
"step": 0.001,
|
||||
"round": False,
|
||||
"tooltip": "Attempts to desaturate the latent by injecting the average across channels (controlled by channel_correction). Applied after mix.",
|
||||
},
|
||||
),
|
||||
"channel_correlation": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "1, 1, 1, 1, 1, 1",
|
||||
"multiline": False,
|
||||
"dynamicPrompts": False,
|
||||
"tooltip": "Comma-separated list of channel correlation strengths.",
|
||||
},
|
||||
),
|
||||
"preview": (
|
||||
("none", "no_mix", "mix"),
|
||||
{
|
||||
"tooltip": "When enabled, displays a preview of the filter shape and a sample of noise. Mix - previews noise after mix is applied. no_mix - only previews the filtered noise.",
|
||||
},
|
||||
),
|
||||
}
|
||||
return result
|
||||
INPUT_TYPES = (
|
||||
NoiseChainInputTypes()
|
||||
.req_bool_time_brownian(
|
||||
tooltip="Controls whether brownian noise is used when mix isn't 1.0.",
|
||||
)
|
||||
.req_float_alpha(
|
||||
default=0.0,
|
||||
min=-5.0,
|
||||
max=5.0,
|
||||
tooltip="Values above 0 will amplify low frequencies, negative values will amplify high frequencies.",
|
||||
)
|
||||
.req_float_max_freq(
|
||||
default=0.7071,
|
||||
min=0.0,
|
||||
max=0.7071,
|
||||
tooltip="Maximum frequency to pass through the filter.",
|
||||
)
|
||||
.req_float_min_freq(
|
||||
default=0.0,
|
||||
min=0.0,
|
||||
max=0.7071,
|
||||
tooltip="Minimum frequency to pass through the filter.",
|
||||
)
|
||||
.req_float_stretch(
|
||||
default=1.0,
|
||||
min=0.01,
|
||||
max=100.0,
|
||||
tooltip="Stretches the filter's shape by the specified factor.",
|
||||
)
|
||||
.req_float_rotate(
|
||||
default=0.0,
|
||||
min=-90.0,
|
||||
max=90.0,
|
||||
step=5.0,
|
||||
tooltip="Rotates the filter.",
|
||||
)
|
||||
.req_float_pnorm(
|
||||
default=2.0,
|
||||
min=0.125,
|
||||
max=100.0,
|
||||
step=0.1,
|
||||
tooltip="Factor used for cushioning the band-pass region.",
|
||||
)
|
||||
.req_floatpct_mix(
|
||||
default=1.0,
|
||||
tooltip="Controls the ratio of filtered noise. For example, 0.75 means 75% noise with the filter effects applied, 25% raw noise.",
|
||||
)
|
||||
.req_float_common_mode(
|
||||
default=0.0,
|
||||
min=-100.0,
|
||||
max=100.0,
|
||||
tooltip="Attempts to desaturate the latent by injecting the average across channels (controlled by channel_correction). Applied after mix.",
|
||||
)
|
||||
.req_string_channel_correlation(
|
||||
default="1, 1, 1, 1, 1, 1",
|
||||
tooltip="Comma-separated list of channel correlation strengths.",
|
||||
)
|
||||
.req_field_preview(
|
||||
("none", "no_mix", "mix"),
|
||||
default="none",
|
||||
tooltip="When enabled, displays a preview of the filter shape and a sample of noise. Mix - previews noise after mix is applied. no_mix - only previews the filtered noise.",
|
||||
)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_item_class(cls):
|
||||
@@ -664,7 +629,7 @@ class SonarPowerNoiseNode(SonarCustomNoiseNodeBase):
|
||||
def go(
|
||||
self,
|
||||
preview="none",
|
||||
**kwargs,
|
||||
**kwargs: Any,
|
||||
):
|
||||
result = super().go(**kwargs)
|
||||
if preview == "none":
|
||||
@@ -680,7 +645,7 @@ class SonarPowerFilterNoiseNode(SonarPowerNoiseNode, SonarNormalizeNoiseNodeMixi
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
result = super().INPUT_TYPES(include_rescale=False, include_chain=False)
|
||||
result = super().INPUT_TYPES()
|
||||
for k in (
|
||||
"min_freq",
|
||||
"max_freq",
|
||||
@@ -749,7 +714,7 @@ class SonarPowerFilterNoiseNode(SonarPowerNoiseNode, SonarNormalizeNoiseNodeMixi
|
||||
normalize_noise,
|
||||
normalize_result,
|
||||
preview="none",
|
||||
**kwargs: dict,
|
||||
**kwargs: Any,
|
||||
):
|
||||
return super().go(
|
||||
factor=factor,
|
||||
@@ -861,67 +826,42 @@ class SonarPreviewFilterNode:
|
||||
FUNCTION = "go"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"sonar_power_filter": (
|
||||
"SONAR_POWER_FILTER",
|
||||
{
|
||||
"tooltip": "Power Filter to preview.",
|
||||
},
|
||||
),
|
||||
"filter_gain": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1 / 3,
|
||||
"min": 0.0,
|
||||
"max": 1000000.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Gain factor applied to the filter part of the preview.",
|
||||
},
|
||||
),
|
||||
"kernel_gain": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1 / 3,
|
||||
"min": 0.0,
|
||||
"max": 1000000.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Gain factor applied to the kernel part of the preview.",
|
||||
},
|
||||
),
|
||||
"norm_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.1,
|
||||
"round": False,
|
||||
"tooltip": "Normalization factor applied to the filter before previewing. 1.0 means 100% normalized.",
|
||||
},
|
||||
),
|
||||
"preview_size": (
|
||||
(
|
||||
"128x128",
|
||||
"256x256",
|
||||
"384x256",
|
||||
"256x384",
|
||||
"768x512",
|
||||
"512x768",
|
||||
"768x768",
|
||||
"128x127",
|
||||
"127x128",
|
||||
),
|
||||
{
|
||||
"tooltip": "Controls the size of the generated preview. Note: Sizes are in latent pixels. For most models, one latent pixel equals eight pixels",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
INPUT_TYPES = (
|
||||
SonarInputTypes()
|
||||
.req_field_sonar_power_filter(
|
||||
"SONAR_POWER_FILTER",
|
||||
tooltip="Power Filter to preview.",
|
||||
)
|
||||
.req_float_filter_gain(
|
||||
default=1 / 3,
|
||||
min=0.0,
|
||||
tooltip="Gain factor applied to the filter part of the preview.",
|
||||
)
|
||||
.req_float_kernel_gain(
|
||||
default=1 / 3,
|
||||
min=0.0,
|
||||
tooltip="Gain factor applied to the kernel part of the preview.",
|
||||
)
|
||||
.req_floatpct_norm_factor(
|
||||
default=1.0,
|
||||
tooltip="Normalization factor applied to the filter before previewing. 1.0 means 100% normalized.",
|
||||
)
|
||||
.req_field_preview_size(
|
||||
(
|
||||
"128x128",
|
||||
"256x256",
|
||||
"384x256",
|
||||
"256x384",
|
||||
"768x512",
|
||||
"512x768",
|
||||
"768x768",
|
||||
"128x127",
|
||||
"127x128",
|
||||
),
|
||||
default="128x128",
|
||||
tooltip="Controls the size of the generated preview. Note: Sizes are in latent pixels. For most models, one latent pixel equals eight pixels",
|
||||
)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def go(
|
||||
@@ -944,3 +884,11 @@ class SonarPreviewFilterNode:
|
||||
),
|
||||
(filt,),
|
||||
)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SonarPowerNoise": SonarPowerNoiseNode,
|
||||
"SonarPowerFilterNoise": SonarPowerFilterNoiseNode,
|
||||
"SonarPowerFilter": SonarPowerFilterNode,
|
||||
"SonarPreviewFilter": SonarPreviewFilterNode,
|
||||
}
|
||||
+1305
-200
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,55 @@
|
||||
from .automata_noise_generator import AutomataNoiseGenerator
|
||||
from .base import MixedNoiseGenerator, NoiseError, NoiseType
|
||||
from .collatz_noise_generator import CollatzNoiseGenerator
|
||||
from .distro_noise_generator import DistroNoiseGenerator
|
||||
from .novelty_filtered_noise import NoveltyFilteredNoiseGenerator
|
||||
from .scatternet_filtered_noise_generator import ScatternetFilteredNoiseGenerator
|
||||
from .simple_noise_generators import (
|
||||
BrownianNoiseGenerator,
|
||||
GaussianNoiseGenerator,
|
||||
GreenTestNoiseGenerator,
|
||||
HighresPyramidNoiseGenerator,
|
||||
LaplacianNoiseGenerator,
|
||||
OneFNoiseGenerator,
|
||||
PerlinOldNoiseGenerator,
|
||||
PinkOldNoiseGenerator,
|
||||
PowerLawNoiseGenerator,
|
||||
PowerOldNoiseGenerator,
|
||||
PyramidNoiseGenerator,
|
||||
PyramidOldNoiseGenerator,
|
||||
StudentTNoiseGenerator,
|
||||
UniformNoiseGenerator,
|
||||
)
|
||||
from .simulation_noise_generator import SimulationNoiseGenerator
|
||||
from .voronoi_noise_generator import VoronoiNoiseGenerator
|
||||
from .wavelet_filtered_noise_generator import WaveletFilteredNoiseGenerator
|
||||
from .wavelet_noise_generator import WaveletNoiseGenerator
|
||||
|
||||
__all__ = (
|
||||
"AutomataNoiseGenerator",
|
||||
"BrownianNoiseGenerator",
|
||||
"CollatzNoiseGenerator",
|
||||
"DistroNoiseGenerator",
|
||||
"GaussianNoiseGenerator",
|
||||
"GreenTestNoiseGenerator",
|
||||
"HighresPyramidNoiseGenerator",
|
||||
"LaplacianNoiseGenerator",
|
||||
"MixedNoiseGenerator",
|
||||
"NoiseError",
|
||||
"NoiseType",
|
||||
"NoveltyFilteredNoiseGenerator",
|
||||
"OneFNoiseGenerator",
|
||||
"PerlinOldNoiseGenerator",
|
||||
"PinkOldNoiseGenerator",
|
||||
"PowerLawNoiseGenerator",
|
||||
"PowerOldNoiseGenerator",
|
||||
"PyramidNoiseGenerator",
|
||||
"PyramidOldNoiseGenerator",
|
||||
"ScatternetFilteredNoiseGenerator",
|
||||
"SimulationNoiseGenerator",
|
||||
"StudentTNoiseGenerator",
|
||||
"UniformNoiseGenerator",
|
||||
"VoronoiNoiseGenerator",
|
||||
"WaveletFilteredNoiseGenerator",
|
||||
"WaveletNoiseGenerator",
|
||||
)
|
||||
@@ -0,0 +1,325 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import torch
|
||||
from comfy.model_management import throw_exception_if_processing_interrupted
|
||||
from tqdm import trange
|
||||
|
||||
from .base import NoiseGenerator
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
|
||||
# Analytic extension of the Collatz conjecture for floating point numbers.
|
||||
def continuous_collatz(x: torch.Tensor) -> torch.Tensor:
|
||||
# f(x) = 1/4 * (2 + 7x - (2 + 5x)*cos(pi*x))
|
||||
cos_term = (x * torch.pi).cos_()
|
||||
return x.mul(7).add_(2).sub_(x.mul(5).add_(2).mul_(cos_term)).mul_(0.25)
|
||||
|
||||
|
||||
def generate_spatial_collatz_noise(
|
||||
batch: int,
|
||||
channels: int,
|
||||
height: int,
|
||||
width: int,
|
||||
depth: int = None, # Optional 3D depth
|
||||
steps: int = 20,
|
||||
num_seeds: int = 15,
|
||||
device: str = "cpu",
|
||||
):
|
||||
is_3d = depth is not None
|
||||
|
||||
# 1. Initialize grid
|
||||
if is_3d:
|
||||
grid = torch.zeros((batch, channels, depth, height, width), device=device)
|
||||
norm_dims = [2, 3, 4]
|
||||
else:
|
||||
grid = torch.zeros((batch, channels, height, width), device=device)
|
||||
norm_dims = [2, 3]
|
||||
|
||||
# 2. Plant float/negative "seeds"
|
||||
for b in range(batch):
|
||||
for c in range(channels):
|
||||
seed_y = torch.randint(0, height, (num_seeds,))
|
||||
seed_x = torch.randint(0, width, (num_seeds,))
|
||||
|
||||
# Using random floats from -1000 to 1000
|
||||
seed_vals = (
|
||||
torch.rand((num_seeds,), dtype=torch.float32, device=device) * 2000.0
|
||||
) - 1000.0
|
||||
|
||||
if is_3d:
|
||||
seed_z = torch.randint(0, depth, (num_seeds,))
|
||||
grid[b, c, seed_z, seed_y, seed_x] = seed_vals
|
||||
else:
|
||||
grid[b, c, seed_y, seed_x] = seed_vals
|
||||
|
||||
# 3. Create spatial diffusion kernel
|
||||
if is_3d:
|
||||
# Create a 3x3x3 blurring kernel using outer products
|
||||
k1d = torch.tensor([1.0, 2.0, 1.0], device=device)
|
||||
kernel = (k1d.view(3, 1, 1) * k1d.view(1, 3, 1) * k1d.view(1, 1, 3)) / 64.0
|
||||
kernel = kernel.view(1, 1, 3, 3, 3).repeat(channels, 1, 1, 1, 1)
|
||||
conv_fn = F.conv3d
|
||||
else:
|
||||
# Create a 3x3 blurring kernel
|
||||
k1d = torch.tensor([1.0, 2.0, 1.0], device=device)
|
||||
kernel = (k1d.view(3, 1) * k1d.view(1, 3)) / 16.0
|
||||
kernel = kernel.view(1, 1, 3, 3).repeat(channels, 1, 1, 1)
|
||||
conv_fn = F.conv2d
|
||||
|
||||
# 4. Evolve the grid
|
||||
for _ in trange(steps, desc="Automata", miniter=25):
|
||||
# A. Spatial diffusion (spread values into neighboring dimensions)
|
||||
grid = conv_fn(grid, kernel, padding=1, groups=channels)
|
||||
|
||||
# B. Apply Collatz activation
|
||||
grid = continuous_collatz(grid)
|
||||
|
||||
# C. Reset rule: inject new seeds if elements get trapped in low magnitude cycles
|
||||
trapped_mask = grid.abs() <= 1.5
|
||||
if trapped_mask.any():
|
||||
new_seeds = (torch.rand_like(grid) * 200.0) - 100.0
|
||||
grid = torch.where(trapped_mask, new_seeds, grid)
|
||||
|
||||
# D. Internal Instance Normalization to tame the math
|
||||
mean = grid.mean(dim=norm_dims, keepdim=True)
|
||||
std = grid.std(dim=norm_dims, keepdim=True) + 1e-5
|
||||
grid = (grid - mean) / std
|
||||
|
||||
# 5. Final Output Normalization
|
||||
mean = grid.mean(dim=norm_dims, keepdim=True)
|
||||
std = grid.std(dim=norm_dims, keepdim=True) + 1e-5
|
||||
|
||||
return (grid - mean) / std
|
||||
|
||||
|
||||
class AutomataNoiseGenerator(NoiseGenerator):
|
||||
name = "automata"
|
||||
|
||||
blend_function: Callable | None = None
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
# Evolution mode
|
||||
# collatz, logistic, sawtooth, lenia, roll
|
||||
"evolution_mode": "collatz",
|
||||
# blur, laplacian, crystal
|
||||
"spread_mode": "blur",
|
||||
"steps": 20,
|
||||
"spread_substeps": 3,
|
||||
"num_seeds": 10,
|
||||
"depth": 10,
|
||||
"trapped_threshold": 1.5,
|
||||
"trapped_interval": 1,
|
||||
# Controls behavior for trapped elements.
|
||||
# new - new seed, reset - original seed, mean - replace with mean
|
||||
"trapped_mode": "reset",
|
||||
"range_negative": -100.0,
|
||||
"range_positive": 100.0,
|
||||
# Absolute value.
|
||||
"seed_minimum": 1.5,
|
||||
"noise_sampler_factory": None,
|
||||
}
|
||||
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
if not (self.height and self.width):
|
||||
raise ValueError("Unsupported shape")
|
||||
self.grid = self.grid_orig = None
|
||||
self.noise_chunk = None
|
||||
self.current_depth = 0
|
||||
|
||||
def create_grid(self) -> None:
|
||||
batch, channels = self.batch, self.channels
|
||||
height, width, depth = self.height, self.width, self.depth
|
||||
num_seeds = self.num_seeds
|
||||
is_3d = depth > 0
|
||||
device, dtype = self.gen_device, self.dtype
|
||||
total_seeds = batch * channels * num_seeds
|
||||
|
||||
# Create flat arrays of coordinates for every single seed
|
||||
batch_idx = (
|
||||
torch.arange(batch, device=device)
|
||||
.view(-1, 1, 1)
|
||||
.expand(batch, channels, num_seeds)
|
||||
.flatten()
|
||||
)
|
||||
chan_idx = (
|
||||
torch.arange(channels, device=device)
|
||||
.view(1, -1, 1)
|
||||
.expand(batch, channels, num_seeds)
|
||||
.flatten()
|
||||
)
|
||||
y_idx = torch.randint(
|
||||
0,
|
||||
height,
|
||||
(total_seeds,),
|
||||
device=device,
|
||||
generator=self.generator,
|
||||
)
|
||||
x_idx = torch.randint(
|
||||
0,
|
||||
width,
|
||||
(total_seeds,),
|
||||
device=device,
|
||||
generator=self.generator,
|
||||
)
|
||||
|
||||
if is_3d:
|
||||
grid = torch.zeros(
|
||||
(batch, channels, depth, height, width),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
else:
|
||||
grid = torch.zeros(
|
||||
(batch, channels, height, width),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
seed_vals = (
|
||||
torch.rand(
|
||||
(total_seeds,), dtype=dtype, device=device, generator=self.generator
|
||||
)
|
||||
* 2000.0
|
||||
) - 1000.0
|
||||
seed_vals = seed_vals.abs().clamp_min(self.seed_minimum).copysign(seed_vals)
|
||||
|
||||
if is_3d:
|
||||
z_idx = torch.randint(0, depth, (total_seeds,), device=device)
|
||||
grid[batch_idx, chan_idx, z_idx, y_idx, x_idx] = seed_vals
|
||||
else:
|
||||
grid[batch_idx, chan_idx, y_idx, x_idx] = seed_vals
|
||||
self.grid = grid
|
||||
self.initial_grid = grid.clone()
|
||||
|
||||
def evolve_step(self, grid: torch.Tensor) -> torch.Tensor:
|
||||
depth = 0 if grid.ndim < 5 else grid.shape[-3]
|
||||
# k1d = torch.tensor([1.0, 2.0, 1.0], device=device)
|
||||
k1d = torch.tensor(
|
||||
[0.1, 1.0, 0.1],
|
||||
device=grid.device,
|
||||
dtype=grid.dtype,
|
||||
)
|
||||
if depth > 0:
|
||||
# Create a 3x3x3 blurring kernel using outer products
|
||||
kernel = k1d.view(3, 1, 1) * k1d.view(1, 3, 1) * k1d.view(1, 1, 3)
|
||||
kernel /= kernel.sum()
|
||||
kernel = kernel.view(1, 1, 3, 3, 3).repeat(self.channels, 1, 1, 1, 1)
|
||||
else:
|
||||
# Create a 3x3 blurring kernel
|
||||
kernel = k1d.view(3, 1) * k1d.view(1, 3)
|
||||
kernel /= kernel.sum()
|
||||
kernel = kernel.view(1, 1, 3, 3).repeat(self.channels, 1, 1, 1)
|
||||
op = partial(
|
||||
F.conv3d if depth > 0 else F.conv2d,
|
||||
weight=kernel,
|
||||
padding=1,
|
||||
groups=self.channels,
|
||||
)
|
||||
# op = partial(F.max_pool3d if depth > 0 else F.max_pool2d, kernel_size=3, stride=1, padding=1)
|
||||
for _ in range(self.spread_substeps):
|
||||
grid = grid.lerp(op(grid), 1.0)
|
||||
# grid = F.max_pool3d(grid, 3, stride=1, padding=1)
|
||||
# grid = conv_fn(grid, kernel, padding=1, groups=self.channels)
|
||||
grid = continuous_collatz(grid)
|
||||
return grid
|
||||
# return continuous_collatz(grid)
|
||||
|
||||
def handle_trapped(
|
||||
self,
|
||||
*,
|
||||
grid: torch.Tensor,
|
||||
orig_grid: torch.Tensor,
|
||||
grid_prev: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
if self.trapped_threshold == 0:
|
||||
return grid
|
||||
mask = grid.abs() < self.trapped_threshold
|
||||
if grid_prev is not None:
|
||||
mask &= grid_prev.abs() >= self.trapped_threshold
|
||||
if not torch.any(mask):
|
||||
return grid
|
||||
|
||||
new_seeds = (torch.rand_like(grid) * 2000.0) - 1000.0
|
||||
return torch.where(mask, new_seeds, grid)
|
||||
# return torch.where(mask, orig_grid, grid) if torch.any(mask) else grid
|
||||
|
||||
def handle_norm(self, grid: torch.Tensor) -> torch.Tensor:
|
||||
return grid.clamp(-10000.0, 10000.0)
|
||||
dims = tuple(range(2, grid.ndim))
|
||||
gn = grid.clone()
|
||||
gn /= gn.std(dim=dims, keepdim=True).clamp_min_(1e-06)
|
||||
return grid.lerp(gn, grid.abs().div_(10000.0).clamp_max_(1.0))
|
||||
# mask = grid.abs() > 100.0
|
||||
# return torch.where(mask, grid.lerp(gn, 0.5), grid)
|
||||
# gn = grid - grid.mean(dim=dims, keepdim=True)
|
||||
|
||||
def handle_norm_(self, grid: torch.Tensor) -> torch.Tensor:
|
||||
# return (grid.abs() % 1000000.0).copysign_(grid)
|
||||
# mask = grid.abs() > 40000.0
|
||||
# return torch.where(
|
||||
# mask,
|
||||
# (grid.cos() * 1000.0).abs().clamp_min(1.5).copysign(grid),
|
||||
# grid,
|
||||
# )
|
||||
# return torch.where(mask, (grid.abs() % 2000.0).copysign(grid), grid)
|
||||
# new_seeds = (torch.rand_like(grid) * 2000.0) - 1000.0
|
||||
# return torch.where(mask, new_seeds, grid)
|
||||
# return grid * (~mask).to(grid)
|
||||
mask = grid.abs() > 100000000.0
|
||||
grid = (grid.abs() % 100000000.0).copysign(grid)
|
||||
return grid
|
||||
dims = tuple(range(2, grid.ndim))
|
||||
std = grid.std(dim=dims, keepdim=True)
|
||||
std = std.abs().clamp_min_(1e-08).copysign(std)
|
||||
grid_adj = grid / std
|
||||
grid_adj -= grid_adj.mean(dim=dims, keepdim=True)
|
||||
grid = torch.where(mask, grid_adj, grid)
|
||||
return grid
|
||||
|
||||
def evolve(self):
|
||||
if self.grid is None:
|
||||
self.create_grid()
|
||||
grid = self.grid
|
||||
for i in trange(self.steps, miniters=10, desc="Automata step"):
|
||||
if i > 1 and (i % 5) == 0:
|
||||
throw_exception_if_processing_interrupted()
|
||||
grid_prev = grid
|
||||
grid = self.evolve_step(grid)
|
||||
grid = self.handle_trapped(
|
||||
grid=grid,
|
||||
orig_grid=self.initial_grid,
|
||||
grid_prev=grid_prev,
|
||||
)
|
||||
grid = self.handle_norm(grid)
|
||||
self.grid = grid
|
||||
|
||||
def reset_grid(self):
|
||||
self.grid = self.initial_grid = None
|
||||
self.current_depth = 0
|
||||
|
||||
def generate(self, *args) -> torch.Tensor:
|
||||
if self.grid is None:
|
||||
self.create_grid()
|
||||
self.current_depth = 0
|
||||
self.evolve()
|
||||
if self.grid.ndim < 5:
|
||||
return self.grid.clone()
|
||||
grid = self.grid
|
||||
self.reset_grid()
|
||||
return grid
|
||||
result = self.grid[:, :, self.current_depth, ...].clone()
|
||||
self.current_depth += 1
|
||||
if self.current_depth >= self.grid.shape[-3]:
|
||||
self.reset_grid()
|
||||
return result
|
||||
@@ -0,0 +1,237 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum, auto
|
||||
|
||||
import torch
|
||||
|
||||
from ..utils import (
|
||||
fallback,
|
||||
scale_noise,
|
||||
tensor_to,
|
||||
)
|
||||
|
||||
# ruff: noqa: ANN002, ANN003
|
||||
|
||||
|
||||
class NoiseType(Enum):
|
||||
BROWNIAN = auto()
|
||||
COLLATZ = auto()
|
||||
DISTRO = auto()
|
||||
GAUSSIAN = auto()
|
||||
GREEN_TEST = auto()
|
||||
GREY = auto()
|
||||
HIGHRES_PYRAMID = auto()
|
||||
HIGHRES_PYRAMID_AREA = auto()
|
||||
HIGHRES_PYRAMID_BISLERP = auto()
|
||||
LAPLACIAN = auto()
|
||||
ONEF_GREENISH = auto()
|
||||
ONEF_GREENISH_MIX = auto()
|
||||
ONEF_PINKISH = auto()
|
||||
ONEF_PINKISH_MIX = auto()
|
||||
ONEF_PINKISHGREENISH = auto()
|
||||
PERLIN = auto()
|
||||
PINK_OLD = auto()
|
||||
POWER_OLD = auto()
|
||||
PYRAMID = auto()
|
||||
PYRAMID_AREA = auto()
|
||||
PYRAMID_BISLERP = auto()
|
||||
PYRAMID_DISCOUNT5 = auto()
|
||||
PYRAMID_MIX = auto()
|
||||
PYRAMID_MIX_AREA = auto()
|
||||
PYRAMID_MIX_BISLERP = auto()
|
||||
PYRAMID_OLD = auto()
|
||||
PYRAMID_OLD_AREA = auto()
|
||||
PYRAMID_OLD_BISLERP = auto()
|
||||
RAINBOW_INTENSE = auto()
|
||||
RAINBOW_MILD = auto()
|
||||
STUDENTT = auto()
|
||||
UNIFORM = auto()
|
||||
VELVET = auto()
|
||||
VIOLET = auto()
|
||||
VORONOI_FUZZ = auto()
|
||||
VORONOI_MIX = auto()
|
||||
WAVELET = auto()
|
||||
WHITE = auto()
|
||||
|
||||
@classmethod
|
||||
def get_names(cls, default=GAUSSIAN, skip=None):
|
||||
if default is not None:
|
||||
if isinstance(default, int):
|
||||
default = cls(default)
|
||||
yield default.name.lower()
|
||||
for nt in cls:
|
||||
if nt == default or (skip and nt in skip):
|
||||
continue
|
||||
yield nt.name.lower()
|
||||
|
||||
|
||||
class NoiseError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class NoiseGenerator:
|
||||
name = "unknown"
|
||||
MIN_DIMS = 1
|
||||
MAX_DIMS = 0
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
x,
|
||||
**kwargs,
|
||||
):
|
||||
if x.ndim < self.MIN_DIMS:
|
||||
errstr = f"Noise generator {self.name} requires at least {self.MIN_DIMS} dimension(s) but got input with shape {x.shape}"
|
||||
raise ValueError(errstr)
|
||||
if self.MAX_DIMS > 0 and x.ndim > self.MAX_DIMS:
|
||||
errstr = f"Noise generator {self.name} requires at most {self.MAX_DIMS} dimension(s) but got input with shape {x.shape}"
|
||||
raise ValueError(errstr)
|
||||
params = self.ng_params()
|
||||
kwarg_params = params | kwargs
|
||||
for k in params:
|
||||
setattr(self, k, kwarg_params.pop(k))
|
||||
self.options = kwarg_params
|
||||
self.update_x(x)
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return {
|
||||
"normalized": True,
|
||||
"force_normalize": None,
|
||||
"normalize_dims": None,
|
||||
"cpu": True,
|
||||
"generator": None,
|
||||
}
|
||||
|
||||
def update_x(self, x):
|
||||
self.shape = x.shape
|
||||
self.batch = self.channels = self.frames = self.height = self.width = None
|
||||
if x.ndim >= 2:
|
||||
self.batch, self.channels = x.shape[:2]
|
||||
if x.ndim > 2:
|
||||
self.width = x.shape[-1]
|
||||
if x.ndim > 3:
|
||||
self.height = x.shape[-2]
|
||||
if x.ndim == 5:
|
||||
self.frames = x.shape[-3]
|
||||
self.device = x.device
|
||||
self.gen_device = torch.device("cpu") if self.cpu else self.device
|
||||
self.layout = x.layout
|
||||
self.dtype = x.dtype
|
||||
|
||||
def rand_like(
|
||||
self,
|
||||
*,
|
||||
fun=torch.randn,
|
||||
cpu=None,
|
||||
to_device=True,
|
||||
shape=None,
|
||||
dtype=None,
|
||||
layout=None,
|
||||
device=None,
|
||||
generator=None,
|
||||
):
|
||||
cpu = fallback(cpu, self.cpu)
|
||||
noise = fun(
|
||||
*fallback(shape, self.shape),
|
||||
generator=fallback(generator, self.generator),
|
||||
dtype=fallback(dtype, self.dtype),
|
||||
layout=fallback(layout, self.layout),
|
||||
device=fallback(device, "cpu" if cpu else self.gen_device),
|
||||
)
|
||||
if to_device and noise.device != self.device:
|
||||
noise = tensor_to(noise, self.device)
|
||||
return noise
|
||||
|
||||
def output_hook(self, noise):
|
||||
if noise.device != self.device:
|
||||
noise = tensor_to(noise, self.device)
|
||||
return scale_noise(
|
||||
noise,
|
||||
normalized=self.normalized
|
||||
and (self.force_normalize is None or self.force_normalize is True),
|
||||
normalize_dims=self.normalize_dims,
|
||||
)
|
||||
|
||||
def pre_hook(self):
|
||||
pass
|
||||
|
||||
def generate(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
self.pre_hook()
|
||||
return self.output_hook(self.generate(*args, **kwargs))
|
||||
|
||||
def __str__(self):
|
||||
pretty_params = ", ".join(f"{k}={getattr(self, k)!s}" for k in self.ng_params())
|
||||
return f"<NoiseGenerator({self.name}): device={self.device}, shape={self.shape}, dtype={self.dtype}, {pretty_params}>"
|
||||
|
||||
|
||||
class FramesToChannelsNoiseGenerator(NoiseGenerator):
|
||||
MIN_DIMS = 4
|
||||
MAX_DIMS = 5
|
||||
|
||||
def get_adjusted_shape(self):
|
||||
if self.frames:
|
||||
return (self.batch, self.channels * self.frames, self.height, self.width)
|
||||
return (self.batch, self.channels, self.height, self.width)
|
||||
|
||||
def fix_output_frames(self, noise):
|
||||
if not self.frames:
|
||||
return noise
|
||||
return noise.reshape(
|
||||
self.batch,
|
||||
self.channels,
|
||||
self.frames,
|
||||
self.height,
|
||||
self.width,
|
||||
)
|
||||
|
||||
def rand_like(self, *args, shape=None, **kwargs):
|
||||
noise = super().rand_like(*args, shape=shape, **kwargs)
|
||||
if shape is not None:
|
||||
return noise
|
||||
adjusted_shape = self.get_adjusted_shape()
|
||||
if noise.shape != adjusted_shape:
|
||||
return noise.reshape(*adjusted_shape)
|
||||
return noise
|
||||
|
||||
|
||||
class MixedNoiseGenerator(NoiseGenerator):
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"name": "mixed_noise",
|
||||
"normalized": True,
|
||||
"pass_args": frozenset(("cpu",)),
|
||||
"noise_mix": (),
|
||||
"output_fun": None,
|
||||
}
|
||||
|
||||
def __init__(self, x, *args, **kwargs):
|
||||
min_dim = max_dim = None
|
||||
self.name = kwargs["name"]
|
||||
for item in kwargs["noise_mix"]:
|
||||
ng_class = item[0] if isinstance(item, (tuple, list)) else item
|
||||
cmin, cmax = ng_class.MIN_DIMS, ng_class.MAX_DIMS
|
||||
min_dim = max(min_dim if min_dim is not None else cmin, cmin)
|
||||
max_dim = min(max_dim if max_dim is not None else cmax, cmax)
|
||||
self.MIN_DIMS = min_dim
|
||||
self.MAX_DIMS = max_dim
|
||||
super().__init__(x, *args, **kwargs)
|
||||
ng_list = []
|
||||
for ng_class, ng_class_kwargs, transform_fun in self.noise_mix:
|
||||
ng_kwargs = {k: v for k, v in kwargs.items() if k in self.pass_args}
|
||||
ng_list.append((ng_class(x, **ng_class_kwargs, **ng_kwargs), transform_fun))
|
||||
self.ng_list = ng_list
|
||||
|
||||
def generate(self, *args):
|
||||
noise = None
|
||||
for ng, transform_fun in self.ng_list:
|
||||
new_noise = ng(*args)
|
||||
if transform_fun is not None:
|
||||
new_noise = transform_fun(new_noise)
|
||||
noise = new_noise if noise is None else noise.add_(new_noise)
|
||||
if self.output_fun is not None:
|
||||
noise = self.output_fun(noise)
|
||||
return noise
|
||||
@@ -0,0 +1,304 @@
|
||||
# ruff: noqa: ANN002
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import TYPE_CHECKING, ClassVar
|
||||
|
||||
import torch
|
||||
from comfy.model_management import throw_exception_if_processing_interrupted
|
||||
|
||||
from .. import utils
|
||||
from ..utils import fallback, normalize_to_scale, tensor_to
|
||||
from .base import NoiseGenerator
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Sequence
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
|
||||
class CollatzNoiseGenerator(NoiseGenerator):
|
||||
name = "collatz"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"adjust_scale": False,
|
||||
"iteration_sign_flipping": True,
|
||||
"chain_length": (1, 1, 2, 2, 3, 3),
|
||||
"iterations": 10,
|
||||
"rmin": -8000.0,
|
||||
"rmax": 8000.0,
|
||||
"flatten": False,
|
||||
"dims": (-1, -1, -2, -2),
|
||||
# values, ratios, mults, adds
|
||||
# seed_x_ratios, seed_x_mults, seed_x_adds
|
||||
# noise_x_ratios, noise_x_mults, noise_x_adds
|
||||
"output_mode": "values",
|
||||
"quantile": 0.5,
|
||||
"quantile_strategy": "clamp",
|
||||
"noise_dtype": torch.float32,
|
||||
"integer_math": True,
|
||||
"even_multiplier": 0.5,
|
||||
"even_addition": 0.0,
|
||||
"odd_multiplier": 3.0,
|
||||
"odd_addition": 1.0,
|
||||
"add_preserves_sign": True,
|
||||
"chain_offset": 5,
|
||||
"break_loops": True,
|
||||
"seed_mode": "default",
|
||||
"seed_noise_sampler": None,
|
||||
"mix_noise_sampler": None,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _get_iter_slices(n_dims, dim, idx, stride) -> tuple:
|
||||
result = [slice(None)] * n_dims
|
||||
result[dim] = slice(idx, None, stride)
|
||||
return tuple(result)
|
||||
|
||||
def _generate_iteration(
|
||||
self,
|
||||
*args,
|
||||
dim: int,
|
||||
chain_length: int,
|
||||
flatten: False,
|
||||
shape=None,
|
||||
):
|
||||
dtype, device = self.dtype, self.device
|
||||
out_shape = shape = fallback(shape, self.shape)
|
||||
if dim >= len(shape):
|
||||
raise ValueError("Requested dimension out of range")
|
||||
rmin, rmax = self.rmin, self.rmax
|
||||
emul, eadd = self.even_multiplier, self.even_addition
|
||||
omul, oadd = self.odd_multiplier, self.odd_addition
|
||||
keepsign = self.add_preserves_sign
|
||||
intmode = self.integer_math
|
||||
rmaxsubmin = rmax - rmin
|
||||
if flatten:
|
||||
shape = torch.Size((*shape[:dim], math.prod(shape[dim:])))
|
||||
size = shape[dim]
|
||||
chain_length = min(size, chain_length)
|
||||
n_chunks = math.ceil(size / chain_length)
|
||||
chain_length += self.chain_offset
|
||||
result_shape = list(shape)
|
||||
chunk_shape = result_shape.copy()
|
||||
result_shape[dim] = chain_length * n_chunks
|
||||
chunk_shape[dim] = n_chunks
|
||||
result = torch.zeros(result_shape, dtype=self.noise_dtype, device=device)
|
||||
adds, muls = result.clone(), result.clone()
|
||||
if self.seed_noise_sampler is not None:
|
||||
orig_noise = self.seed_noise_sampler(*args)[
|
||||
tuple(slice(None, sz) for sz in chunk_shape)
|
||||
].to(result)
|
||||
if flatten:
|
||||
orig_noise = orig_noise.flatten(start_dim=dim)
|
||||
orig_noise = normalize_to_scale(
|
||||
orig_noise[tuple(slice(None, sz) for sz in chunk_shape)],
|
||||
1e-06,
|
||||
1.0,
|
||||
dim=tuple(range(1, len(chunk_shape))),
|
||||
)
|
||||
else:
|
||||
orig_noise = self.rand_like(
|
||||
fun=torch.rand,
|
||||
shape=chunk_shape,
|
||||
dtype=result.dtype,
|
||||
)
|
||||
noise = orig_noise * (rmaxsubmin + 1) + rmin
|
||||
# Derp.
|
||||
noise = torch.where(noise == 0, noise.max() / noise.numel(), noise)
|
||||
if self.seed_mode != "default":
|
||||
noise = torch.where(
|
||||
(noise % 2.0) < 1
|
||||
if self.seed_mode == "force_odd"
|
||||
else (noise % 2.0) >= 1,
|
||||
noise + 1,
|
||||
noise,
|
||||
)
|
||||
if noise.device != self.device:
|
||||
noise = tensor_to(noise, self.device)
|
||||
slice_0 = self._get_iter_slices(result.ndim, dim, 0, chain_length)
|
||||
for chainidx in range(chain_length):
|
||||
if chainidx == 0:
|
||||
muls[slice_0] = 1.0
|
||||
result[slice_0] = noise
|
||||
continue
|
||||
slice_curr = self._get_iter_slices(result.ndim, dim, chainidx, chain_length)
|
||||
slice_prev = self._get_iter_slices(
|
||||
result.ndim,
|
||||
dim,
|
||||
chainidx - 1,
|
||||
chain_length,
|
||||
)
|
||||
prev = result[slice_prev]
|
||||
prev_trunc = utils.trunc_decimals(prev, 2)
|
||||
need_reset = (
|
||||
((prev_trunc >= 1.0) & (prev_trunc < 1.001))
|
||||
| (prev_trunc.abs() < 0.001)
|
||||
if self.break_loops
|
||||
else False
|
||||
)
|
||||
prev_evens = prev % 2 < 1.0
|
||||
prev_adds, prev_muls = adds[slice_prev], muls[slice_prev]
|
||||
muls_next = (
|
||||
torch.where(
|
||||
prev_evens,
|
||||
prev_muls if emul == 1 else prev_muls * emul,
|
||||
prev_muls if omul == 1 else prev_muls * omul,
|
||||
)
|
||||
if emul != 1 or omul != 1
|
||||
else prev_muls
|
||||
)
|
||||
muls[slice_curr] = (
|
||||
torch.where(need_reset, 1.0, muls_next)
|
||||
if need_reset is not False
|
||||
else muls_next
|
||||
)
|
||||
curr_muls = muls[slice_curr]
|
||||
prev_adds_scaled = prev_adds * curr_muls
|
||||
prev_sign = prev.sign() if keepsign else 1.0
|
||||
adds_next = (
|
||||
torch.where(
|
||||
prev_evens,
|
||||
prev_adds_scaled
|
||||
if eadd == 0
|
||||
else prev_adds_scaled + eadd * prev_sign,
|
||||
prev_adds_scaled
|
||||
if oadd == 0
|
||||
else prev_adds_scaled + oadd * prev_sign,
|
||||
)
|
||||
if eadd != 0 or oadd != 0
|
||||
else prev_adds_scaled
|
||||
)
|
||||
adds[slice_curr] = (
|
||||
torch.where(need_reset, 0.0, adds_next)
|
||||
if need_reset is not False
|
||||
else adds_next
|
||||
)
|
||||
curr_adds = adds[slice_curr]
|
||||
result_next = utils.maybe_apply(
|
||||
(noise * curr_muls).add_(curr_adds),
|
||||
intmode,
|
||||
torch.trunc,
|
||||
)
|
||||
result[slice_curr] = (
|
||||
torch.where(need_reset, noise, result_next)
|
||||
if need_reset is not False
|
||||
else result_next
|
||||
)
|
||||
output_slice = tuple(
|
||||
slice(None, sz) for sz in (shape if flatten else out_shape)
|
||||
)
|
||||
return self._iteration_output(
|
||||
*args,
|
||||
result_chains=result,
|
||||
orig_noise=orig_noise,
|
||||
noise=noise,
|
||||
raw_adds=adds,
|
||||
muls=muls,
|
||||
chain_length=chain_length,
|
||||
dim=dim,
|
||||
output_shape=out_shape,
|
||||
output_slice=output_slice,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
def _trim_chain_offset(
|
||||
self,
|
||||
t: torch.Tensor,
|
||||
dim: int,
|
||||
chain_length: int,
|
||||
) -> torch.Tensor:
|
||||
co = self.chain_offset
|
||||
if co < 1:
|
||||
return t
|
||||
chunks = t.split(chain_length, dim)
|
||||
slices = tuple(
|
||||
slice(None) if i != dim else slice(co, None) for i in range(t.ndim)
|
||||
)
|
||||
return torch.cat(
|
||||
tuple(chunk[slices] for chunk in chunks),
|
||||
dim=dim,
|
||||
)
|
||||
|
||||
def _iteration_output(
|
||||
self,
|
||||
*args,
|
||||
result_chains: torch.Tensor,
|
||||
orig_noise: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
raw_adds: torch.Tensor,
|
||||
muls: torch.Tensor,
|
||||
chain_length: int,
|
||||
dim: int,
|
||||
output_shape: Sequence,
|
||||
output_slice: Sequence,
|
||||
dtype: str | torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
omode = self.output_mode
|
||||
quantile = self.quantile
|
||||
noise_exp = noise.repeat_interleave(chain_length, dim)
|
||||
nadds = raw_adds.div_(noise_exp)
|
||||
ratios = result_chains / noise_exp
|
||||
if omode in {"values", "ratios", "seed_x_ratios", "noise_x_ratios"}:
|
||||
out1 = ratios
|
||||
elif omode in {"mults", "seed_x_mults", "noise_x_mults"}:
|
||||
out1 = muls
|
||||
elif omode in {"adds", "seed_x_adds", "noise_x_adds"}:
|
||||
out1 = nadds
|
||||
else:
|
||||
raise ValueError("Bad output mode")
|
||||
out1 = self._trim_chain_offset(out1, dim=dim, chain_length=chain_length)
|
||||
if quantile not in {0, 1}:
|
||||
out1 = utils.quantile_normalize(
|
||||
out1,
|
||||
quantile=quantile,
|
||||
dim=0,
|
||||
strategy=self.quantile_strategy,
|
||||
)
|
||||
out1 = out1[output_slice].reshape(output_shape).to(dtype=dtype)
|
||||
if omode in {"ratios", "mults", "adds"}:
|
||||
return out1
|
||||
if omode in {"values", "seed_x_ratios", "seed_x_mults", "seed_x_adds"}:
|
||||
out2 = orig_noise.repeat_interleave(chain_length - self.chain_offset, dim)
|
||||
elif omode in {"noise_x_ratios", "noise_x_mults", "noise_x_adds"}:
|
||||
out2 = (
|
||||
self.rand_like(dtype=out1.dtype)
|
||||
if self.mix_noise_sampler is None
|
||||
else self.mix_noise_sampler(*args)
|
||||
)
|
||||
out2 = out2[output_slice].reshape(output_shape).to(dtype=dtype)
|
||||
return out2 * out1
|
||||
|
||||
def generate(self, *args):
|
||||
out_dims = len(self.shape)
|
||||
dims = tuple(dim if dim >= 0 else out_dims + dim for dim in self.dims)
|
||||
n_dims, n_chainlens = len(dims), len(self.chain_length)
|
||||
if not all(0 <= d < out_dims for d in dims):
|
||||
raise ValueError("Dimension out of range")
|
||||
dtype, device = self.dtype, self.device
|
||||
result = torch.zeros(self.shape, dtype=dtype, device=device)
|
||||
it_scale = 1.0 / self.iterations
|
||||
for iteration in range(self.iterations):
|
||||
if iteration > 0 and (iteration % 25) == 0:
|
||||
# It's soooo slow!
|
||||
throw_exception_if_processing_interrupted()
|
||||
temp = self._generate_iteration(
|
||||
*args,
|
||||
dim=dims[iteration % n_dims],
|
||||
chain_length=self.chain_length[iteration % n_chainlens],
|
||||
flatten=self.flatten,
|
||||
).mul_(
|
||||
it_scale
|
||||
* (-1 if self.iteration_sign_flipping and (iteration & 1) == 1 else 1),
|
||||
)
|
||||
result += temp
|
||||
if self.adjust_scale:
|
||||
result = normalize_to_scale(
|
||||
result,
|
||||
-1.0,
|
||||
1.0,
|
||||
dim=tuple(range(1 if result.ndim < 4 else 2, result.ndim)),
|
||||
)
|
||||
return result
|
||||
@@ -0,0 +1,461 @@
|
||||
# ruff: noqa: ANN002, ANN003
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from ..utils import quantile_normalize
|
||||
from .base import NoiseGenerator
|
||||
|
||||
|
||||
class DistroNoiseGenerator(NoiseGenerator):
|
||||
name = "distro"
|
||||
|
||||
simple_distros = frozenset((
|
||||
"cauchy",
|
||||
"exponential",
|
||||
"geometric",
|
||||
"log_normal",
|
||||
"normal",
|
||||
))
|
||||
|
||||
def __init__(self, x, *args, **kwargs):
|
||||
super().__init__(x, *args, **kwargs)
|
||||
if self.distro not in self.distro_params():
|
||||
raise ValueError("Bad distro")
|
||||
|
||||
_distro_params = None
|
||||
|
||||
@classmethod
|
||||
def distro_params(cls):
|
||||
if cls._distro_params is not None:
|
||||
return cls._distro_params
|
||||
td = torch.distributions
|
||||
tt = torch.Tensor
|
||||
cls._distro_params = {
|
||||
# Simple
|
||||
"exponential": (
|
||||
tt.exponential_,
|
||||
{
|
||||
"lambd": {
|
||||
"default": 1.0,
|
||||
},
|
||||
},
|
||||
),
|
||||
"cauchy": (
|
||||
tt.cauchy_,
|
||||
{
|
||||
"median": {
|
||||
"default": "0.0",
|
||||
},
|
||||
"sigma": {
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
},
|
||||
},
|
||||
),
|
||||
"geometric": (
|
||||
tt.geometric_,
|
||||
{
|
||||
"p": {
|
||||
"default": 0.25,
|
||||
},
|
||||
},
|
||||
),
|
||||
"log_normal": (
|
||||
tt.log_normal_,
|
||||
{
|
||||
"mean": {
|
||||
"default": 1.0,
|
||||
},
|
||||
"std": {
|
||||
"default": 2.0,
|
||||
},
|
||||
},
|
||||
),
|
||||
"normal": (
|
||||
tt.normal_,
|
||||
{
|
||||
"mean": {
|
||||
"default": 0.0,
|
||||
},
|
||||
"std": {
|
||||
"default": 1.0,
|
||||
},
|
||||
},
|
||||
),
|
||||
# Complex distros
|
||||
"beta": (
|
||||
td.Beta,
|
||||
{
|
||||
"concentration0": {
|
||||
"default": "0.5",
|
||||
},
|
||||
"concentration1": {
|
||||
"default": "0.5",
|
||||
},
|
||||
},
|
||||
),
|
||||
"continuous_bernoulli": (
|
||||
td.ContinuousBernoulli,
|
||||
{
|
||||
"probs": {
|
||||
"default": "0.5",
|
||||
},
|
||||
},
|
||||
),
|
||||
"dirichlet": (
|
||||
td.Dirichlet,
|
||||
{
|
||||
"concentration": {
|
||||
"default": "0.5 0.5",
|
||||
},
|
||||
},
|
||||
),
|
||||
"fisher_snedecor": (
|
||||
td.FisherSnedecor,
|
||||
{
|
||||
"df1": {
|
||||
"default": "1.0",
|
||||
},
|
||||
"df2": {
|
||||
"default": "2.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"gamma": (
|
||||
td.Gamma,
|
||||
{
|
||||
"concentration": {
|
||||
"default": "1.0",
|
||||
},
|
||||
"rate": {
|
||||
"default": "1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"gumbel": (
|
||||
td.Gumbel,
|
||||
{
|
||||
"loc": {
|
||||
"default": "1.0",
|
||||
},
|
||||
"scale": {
|
||||
"default": "2.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"inverse_gamma": (
|
||||
td.InverseGamma,
|
||||
{
|
||||
"concentration": {
|
||||
"default": "1.0",
|
||||
},
|
||||
"rate": {
|
||||
"default": "1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"kumaraswamy": (
|
||||
td.Kumaraswamy,
|
||||
{
|
||||
"concentration0": {
|
||||
"default": "1.0",
|
||||
},
|
||||
"concentration1": {
|
||||
"default": "1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"laplacian": (
|
||||
td.Laplace,
|
||||
{
|
||||
"loc": {
|
||||
"default": "0.0",
|
||||
},
|
||||
"scale": {
|
||||
"default": "1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"lkjcholesky": (
|
||||
td.LKJCholesky,
|
||||
{
|
||||
"dim": {
|
||||
"_ty": "INT",
|
||||
"default": 3,
|
||||
},
|
||||
"concentration": {
|
||||
"default": "1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"lrmvariate_normal": (
|
||||
lambda loc, cov_factor, cov_diag: td.LowRankMultivariateNormal(
|
||||
loc=loc,
|
||||
cov_factor=cov_factor.reshape(loc.numel(), -1),
|
||||
cov_diag=cov_diag,
|
||||
),
|
||||
{
|
||||
"loc": {
|
||||
"default": "0.0 0.0",
|
||||
},
|
||||
"cov_factor": {
|
||||
"default": "1.0 0.0",
|
||||
},
|
||||
"cov_diag": {
|
||||
"default": "1.0 1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"mvariate_normal": (
|
||||
lambda loc, cov_multiplier=1.0: td.MultivariateNormal(
|
||||
loc=loc,
|
||||
covariance_matrix=torch.eye(
|
||||
loc.numel(),
|
||||
dtype=loc.dtype,
|
||||
device=loc.device,
|
||||
).mul_(cov_multiplier),
|
||||
),
|
||||
{
|
||||
"loc": {
|
||||
"default": "0.0 0.0",
|
||||
},
|
||||
"cov_multiplier": {
|
||||
"default": 1.0,
|
||||
},
|
||||
},
|
||||
),
|
||||
"pareto": (
|
||||
td.Pareto,
|
||||
{
|
||||
"scale": {
|
||||
"default": "1.0",
|
||||
},
|
||||
"alpha": {
|
||||
"default": "1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"poisson": (
|
||||
td.Poisson,
|
||||
{
|
||||
"rate": {
|
||||
"default": "1.5",
|
||||
},
|
||||
},
|
||||
),
|
||||
"relaxed_bernoulli": (
|
||||
td.RelaxedBernoulli,
|
||||
{
|
||||
"temperature": {
|
||||
"default": 0.75,
|
||||
},
|
||||
"probs": {
|
||||
"default": "0.66",
|
||||
},
|
||||
},
|
||||
),
|
||||
"relaxed_onehotcategorical": (
|
||||
td.RelaxedOneHotCategorical,
|
||||
{
|
||||
"temperature": {
|
||||
"default": 1.5,
|
||||
},
|
||||
"probs": {
|
||||
"default": "0.33 0.66",
|
||||
},
|
||||
},
|
||||
),
|
||||
"studentt": (
|
||||
td.StudentT,
|
||||
{
|
||||
"loc": {
|
||||
"default": "0.0",
|
||||
},
|
||||
"scale": {
|
||||
"default": "1.0",
|
||||
},
|
||||
"df": {
|
||||
"default": "1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"uniform": (
|
||||
td.Uniform,
|
||||
{
|
||||
"low": {
|
||||
"default": 0.0,
|
||||
},
|
||||
"high": {
|
||||
"default": 1.0,
|
||||
},
|
||||
},
|
||||
),
|
||||
"vonmises": (
|
||||
td.VonMises,
|
||||
{
|
||||
"loc": {
|
||||
"default": "1.0",
|
||||
},
|
||||
"concentration": {
|
||||
"default": "1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"weibull": (
|
||||
td.Weibull,
|
||||
{
|
||||
"scale": {
|
||||
"default": "1.0",
|
||||
},
|
||||
"concentration": {
|
||||
"default": "1.0",
|
||||
},
|
||||
},
|
||||
),
|
||||
"wishart": (
|
||||
lambda df, cov_size=2, cov_multiplier=1.0: td.Wishart(
|
||||
df=df,
|
||||
covariance_matrix=torch.eye(
|
||||
int(cov_size),
|
||||
dtype=df.dtype,
|
||||
device=df.device,
|
||||
).mul_(cov_multiplier),
|
||||
),
|
||||
{
|
||||
"df": {
|
||||
"default": "2.0",
|
||||
},
|
||||
"cov_size": {
|
||||
"_ty": "INT",
|
||||
"default": 2,
|
||||
},
|
||||
"cov_multiplier": {
|
||||
"default": 1.0,
|
||||
},
|
||||
},
|
||||
),
|
||||
}
|
||||
return cls._distro_params
|
||||
|
||||
_build_params = None
|
||||
|
||||
@classmethod
|
||||
def build_params(cls):
|
||||
if cls._build_params is not None:
|
||||
return cls._build_params
|
||||
cls._build_params = {
|
||||
f"{tykey}_{pkey}": pval
|
||||
for tykey, tyval in cls.distro_params().items()
|
||||
for pkey, pval in tyval[1].items()
|
||||
if not pkey.startswith("_")
|
||||
}
|
||||
return cls._build_params
|
||||
|
||||
_ng_params = None
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
if cls._ng_params is not None:
|
||||
return cls._ng_params
|
||||
dparams = {
|
||||
k: v["default"]
|
||||
for k, v in cls.build_params().items()
|
||||
if not k.startswith("_")
|
||||
}
|
||||
cls._ng_params = (
|
||||
super().ng_params()
|
||||
| {
|
||||
"distro": "normal",
|
||||
"quantile_norm": 0.85,
|
||||
"quantile_norm_flatten": True,
|
||||
"quantile_norm_dim": 1,
|
||||
"quantile_norm_pow": 0.5,
|
||||
"quantile_norm_fac": 1.0,
|
||||
"result_index": "-1",
|
||||
}
|
||||
| dparams
|
||||
)
|
||||
return cls._ng_params
|
||||
|
||||
def norm_output(self, noise):
|
||||
if noise.ndim > len(self.shape):
|
||||
if noise.shape[: len(self.shape)] != self.shape:
|
||||
errstr = f"Unexpected shape when normalizing distro({self.distro}) noise! Output shape={self.shape}, noise shape={noise.shape}, generator dump: {self}"
|
||||
raise RuntimeError(errstr)
|
||||
selfdims = len(self.shape)
|
||||
result_index = self.result_index
|
||||
if not isinstance(result_index, (tuple, list)):
|
||||
result_index = (result_index,)
|
||||
ri_len = len(result_index)
|
||||
if ri_len == 0:
|
||||
raise ValueError("When result_index is a list, it must not be empty")
|
||||
trim_count = 0
|
||||
while noise.ndim > selfdims:
|
||||
idx = result_index[trim_count % ri_len]
|
||||
if idx < 0:
|
||||
idx = noise.shape[-1] + idx
|
||||
noise = noise[..., max(0, min(noise.shape[-1] - 1, idx))]
|
||||
trim_count += 1
|
||||
return (
|
||||
quantile_normalize(
|
||||
noise,
|
||||
quantile=self.quantile_norm,
|
||||
dim=self.quantile_norm_dim,
|
||||
flatten=self.quantile_norm_flatten,
|
||||
nq_fac=self.quantile_norm_fac,
|
||||
pow_fac=self.quantile_norm_pow,
|
||||
)
|
||||
.reshape(self.shape)
|
||||
.contiguous()
|
||||
)
|
||||
|
||||
def distro_param(self, val, *, simple_fun=None):
|
||||
if isinstance(val, torch.Tensor):
|
||||
return simple_fun(val) if simple_fun is not None else val
|
||||
if isinstance(val, str):
|
||||
val = tuple(float(v) for v in val.split(None))
|
||||
if simple_fun is not None:
|
||||
if isinstance(val, (float, int)):
|
||||
return simple_fun(val)
|
||||
if len(val) > 1:
|
||||
raise ValueError("Couldn't return result as float")
|
||||
return simple_fun(val[0])
|
||||
if not isinstance(val, (tuple, list)):
|
||||
val = (val,)
|
||||
return torch.tensor(
|
||||
val,
|
||||
dtype=self.dtype,
|
||||
device=self.gen_device,
|
||||
)
|
||||
|
||||
def get_distro_kwargs(self, distro, ddef, *, simple=False):
|
||||
return {
|
||||
k: self.distro_param(
|
||||
getattr(self, f"{distro}_{k}"),
|
||||
simple_fun=None
|
||||
if not simple and k != "dim"
|
||||
else (int if k == "dim" else float),
|
||||
)
|
||||
for k in ddef
|
||||
}
|
||||
|
||||
def generate(self, *_args):
|
||||
distro = self.distro
|
||||
dfun, ddef = self.distro_params()[distro]
|
||||
is_simple = distro in self.simple_distros
|
||||
dkwargs = self.get_distro_kwargs(distro, ddef, simple=is_simple)
|
||||
if is_simple:
|
||||
noise = torch.empty(
|
||||
*self.shape,
|
||||
device=self.gen_device,
|
||||
dtype=self.dtype,
|
||||
layout=self.layout,
|
||||
)
|
||||
noise = dfun(noise, **dkwargs)
|
||||
else:
|
||||
dobj = dfun(**dkwargs)
|
||||
noise = (
|
||||
dobj.rsample if getattr(dobj, "has_rsample", False) else dobj.sample
|
||||
)(self.shape)
|
||||
return self.norm_output(noise)
|
||||
@@ -0,0 +1,186 @@
|
||||
# ruff: noqa: ANN002, ANN003
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from .base import NoiseGenerator
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
|
||||
def sum_rms_blend(
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
t: torch.Tensor | float = 1.0,
|
||||
*,
|
||||
orig_shape: torch.Size | tuple[int, ...],
|
||||
dims_a: tuple[int, ...] = (1,),
|
||||
dims_b: tuple[int, ...] = (-1, -2),
|
||||
) -> torch.Tensor:
|
||||
rms_a = a / math.prod(orig_shape[d] for d in dims_a) ** 0.5
|
||||
rms_b = b / math.prod(orig_shape[d] for d in dims_b) ** 0.5
|
||||
variance_a = rms_a.pow_(2.0)
|
||||
variance_b = rms_b.pow_(2.0)
|
||||
result = variance_a
|
||||
result += variance_b * t
|
||||
result /= 1.0 + t
|
||||
result **= 0.5
|
||||
return result
|
||||
|
||||
|
||||
def metrics_blend(
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
t: torch.Tensor | float = 1.0,
|
||||
*,
|
||||
orig_shape: torch.Size | tuple[int, ...],
|
||||
dims_a: tuple[int, ...] = (-1, -2),
|
||||
dims_b: tuple[int, ...] = (1,),
|
||||
use_rms: bool = True,
|
||||
rms_power: float = 2.0,
|
||||
) -> torch.Tensor:
|
||||
count_a = math.prod(orig_shape[d] for d in dims_a)
|
||||
count_b = math.prod(orig_shape[d] for d in dims_b)
|
||||
denom_a = count_a ** (1 / rms_power) if use_rms else count_a
|
||||
denom_b = count_b ** (1 / rms_power) if use_rms else count_b
|
||||
curr_a = a / denom_a
|
||||
curr_b = b / denom_b
|
||||
if use_rms:
|
||||
curr_a **= rms_power
|
||||
curr_b = curr_b.pow_(rms_power) * t
|
||||
result = curr_b.add_(curr_a)
|
||||
result /= 1.0 + t
|
||||
return result.pow_(1.0 / rms_power) if use_rms else result
|
||||
|
||||
|
||||
class NoveltyFilteredNoiseGenerator(NoiseGenerator):
|
||||
name = "novelty"
|
||||
|
||||
initial_noise_state: torch.Tensor | None = None
|
||||
noise_state: torch.Tensor | None = None
|
||||
blend_function: Callable | None = None
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"skip_initial": 1,
|
||||
"iters_per_call": 1,
|
||||
"dim_groups": ((1,), (-1, -2)),
|
||||
"blend_ratio": 1.0,
|
||||
"blend_function": None,
|
||||
"update_blend_ratio": 1.0,
|
||||
"update_blend_function": None,
|
||||
"noise_sampler": None,
|
||||
}
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
if self.blend_function is None:
|
||||
raise ValueError("Missing blend function!")
|
||||
|
||||
def generate(self, *args) -> torch.Tensor:
|
||||
ng = (
|
||||
partial(self.noise_sampler, *args) if self.noise_sampler else self.rand_like
|
||||
)
|
||||
noise_state = self.noise_state
|
||||
had_state = self.noise_state is not None
|
||||
it_counter = 0 if had_state else 0 - self.skip_initial
|
||||
its_call = max(1, self.iters_per_call)
|
||||
bf = self.blend_function
|
||||
blend_ratio = self.blend_ratio
|
||||
update_blend_ratio = self.update_blend_ratio
|
||||
ubf = self.update_blend_function
|
||||
if ubf is None:
|
||||
# Linear weighted average
|
||||
def ubf(a: torch.Tensor, b: torch.Tensor, t: float) -> torch.Tensor:
|
||||
return (b * t).add_(a).div_(1.0 + abs(t))
|
||||
|
||||
call_initial_noise = None
|
||||
curr_noise = None
|
||||
call_noise_state = None
|
||||
while it_counter < its_call:
|
||||
if noise_state is None:
|
||||
noise_state = ng()
|
||||
self.initial_noise_state = noise_state.clone()
|
||||
self.noise_state = noise_state.clone()
|
||||
continue
|
||||
curr_noise = ng()
|
||||
if call_initial_noise is None:
|
||||
call_initial_noise = curr_noise.clone()
|
||||
seen = {id(curr_noise)}
|
||||
for ortho_target in (
|
||||
self.initial_noise_state,
|
||||
call_initial_noise if it_counter > 0 else None,
|
||||
noise_state,
|
||||
):
|
||||
tid = id(ortho_target)
|
||||
if ortho_target is None or tid in seen:
|
||||
continue
|
||||
seen.add(tid)
|
||||
curr_noise = bf(ortho_target, curr_noise, blend_ratio)
|
||||
curr_noise -= ortho_target
|
||||
it_counter += 1
|
||||
if it_counter < 1:
|
||||
self.noise_state = curr_noise.clone()
|
||||
noise_state = curr_noise
|
||||
continue
|
||||
if call_noise_state is None:
|
||||
call_noise_state = curr_noise
|
||||
else:
|
||||
call_noise_state = ubf(call_noise_state, curr_noise, update_blend_ratio)
|
||||
if call_noise_state is None:
|
||||
raise RuntimeError("Unexpected unpopulated call_noise_state!")
|
||||
# self.noise_state = call_noise_state.clone()
|
||||
self.noise_state = ubf(noise_state, call_noise_state, update_blend_ratio)
|
||||
return call_noise_state
|
||||
|
||||
# def generate(self, *args) -> torch.Tensor:
|
||||
# ng = (
|
||||
# partial(self.noise_sampler, *args) if self.noise_sampler else self.rand_like
|
||||
# )
|
||||
# noise_state = self.noise_state
|
||||
# had_state = self.noise_state is not None
|
||||
# it_counter = 0 if had_state else 0 - self.skip_initial
|
||||
# its_call = self.iters_per_call
|
||||
# bf = self.blend_function
|
||||
# blend_ratio = self.blend_ratio
|
||||
# update_blend_ratio = self.update_blend_ratio
|
||||
# ubf = self.update_blend_function
|
||||
# if ubf is None or True:
|
||||
# # Linear weighted average
|
||||
# def ubf(a: torch.Tensor, b: torch.Tensor, t: float) -> torch.Tensor:
|
||||
# return (b * t).add_(a).div_(1.0 + abs(t))
|
||||
|
||||
# call_initial_noise = None
|
||||
# while it_counter < its_call:
|
||||
# if noise_state is None:
|
||||
# noise_state = ng()
|
||||
# self.initial_noise_state = noise_state.clone()
|
||||
# continue
|
||||
# curr_noise = ng()
|
||||
# it_counter += 1
|
||||
# if it_counter < 1:
|
||||
# noise_state = curr_noise
|
||||
# continue
|
||||
# if call_initial_noise is None and self.iters_per_call > 1:
|
||||
# call_initial_noise = curr_noise.clone()
|
||||
# for ortho_target in (
|
||||
# self.initial_noise_state,
|
||||
# call_initial_noise if it_counter > 0 else None,
|
||||
# noise_state,
|
||||
# ):
|
||||
# if ortho_target is None:
|
||||
# continue
|
||||
# curr_noise = bf(ortho_target, curr_noise, blend_ratio)
|
||||
# curr_noise -= ortho_target
|
||||
# # curr_noise = bf(noise_state, ng(), blend_ratio).sub_(noise_state)
|
||||
# noise_state = ubf(noise_state, curr_noise, update_blend_ratio)
|
||||
# self.noise_state = noise_state.clone()
|
||||
# return noise_state
|
||||
@@ -0,0 +1,173 @@
|
||||
# ruff: noqa: ANN002, ANN003
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
from .. import utils
|
||||
from ..wavelet_functions import ptwav
|
||||
from .base import FramesToChannelsNoiseGenerator
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
|
||||
class ScatternetFilteredNoiseGenerator(FramesToChannelsNoiseGenerator):
|
||||
name = "scatternetfilter"
|
||||
MIN_DIMS = 4
|
||||
MAX_DIMS = 4
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
if ptwav is None:
|
||||
raise RuntimeError(
|
||||
"Scatternet noise requires the pytorch_wavelets package to be installed in your Python environment",
|
||||
)
|
||||
super().__init__(*args, **kwargs)
|
||||
if self.output_mode not in {
|
||||
"channels",
|
||||
"channels_adjusted",
|
||||
"channels_scaled",
|
||||
"flat",
|
||||
"flat_adjusted",
|
||||
"flat_scaled",
|
||||
}:
|
||||
raise ValueError("Bad output mode")
|
||||
|
||||
scatkwargs = {
|
||||
"mode": self.mode,
|
||||
"biort": "near_sym_b_bp" if self.use_symmetric_filter else self.biort,
|
||||
}
|
||||
if self.scatternet_order == 2:
|
||||
scatkwargs["qshift"] = (
|
||||
"qshift_b_bp" if self.use_symmetric_filter else self.qshift
|
||||
)
|
||||
self.scatternet = ptwav.ScatLayerj2(**scatkwargs)
|
||||
elif self.scatternet_order == 1:
|
||||
self.scatternet = ptwav.ScatLayer(**scatkwargs)
|
||||
else:
|
||||
self.scatternet = torch.nn.Sequential(
|
||||
*(
|
||||
ptwav.ScatLayer(**scatkwargs)
|
||||
for _ in range(abs(self.scatternet_order))
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"mode": "symmetric",
|
||||
"magbias": 1e-02,
|
||||
"use_symmetric_filter": False,
|
||||
"biort": "near_sym_a",
|
||||
"qshift": "qshift_a",
|
||||
"output_offset": 0.0,
|
||||
"scatternet_order": 1,
|
||||
"per_channel_scatternet": False,
|
||||
"output_mode": "channels_adjusted",
|
||||
# If None, uses probselect when available, otherwise bilinear.
|
||||
"upscale_mode": None,
|
||||
"noise_sampler": None,
|
||||
}
|
||||
|
||||
def _fix_shape(self, noise, adjusted_shape):
|
||||
if self.frames:
|
||||
noise = noise.reshape(
|
||||
self.batch,
|
||||
self.channels * self.frames,
|
||||
self.height,
|
||||
self.width,
|
||||
)
|
||||
elif noise.shape != adjusted_shape:
|
||||
noise = noise.reshape(*adjusted_shape)
|
||||
return noise
|
||||
|
||||
def generate(self, *args):
|
||||
adjusted_shape = self.get_adjusted_shape()
|
||||
scaled = self.output_mode.endswith("_scaled")
|
||||
adjusted = scaled or self.output_mode.endswith("_adjusted")
|
||||
order = abs(self.scatternet_order)
|
||||
order_spatial_compensation = 2**order
|
||||
output_mode = (
|
||||
self.output_mode.split("_", 1)[0] if adjusted else self.output_mode
|
||||
)
|
||||
spatial_compensation = 1 if adjusted else order_spatial_compensation
|
||||
if self.noise_sampler is None:
|
||||
temp_shape = (
|
||||
(
|
||||
*adjusted_shape[:2],
|
||||
adjusted_shape[-2] * spatial_compensation,
|
||||
adjusted_shape[-1] * spatial_compensation,
|
||||
)
|
||||
if spatial_compensation != 1
|
||||
else adjusted_shape
|
||||
)
|
||||
noise = self.rand_like(shape=temp_shape)
|
||||
else:
|
||||
noise = self.noise_sampler(*args)
|
||||
if scaled:
|
||||
upscale_mode = self.upscale_mode
|
||||
if upscale_mode is None:
|
||||
upscale_mode = (
|
||||
"probselect"
|
||||
if "probselect" in utils.UPSCALE_METHODS
|
||||
else "bilinear"
|
||||
)
|
||||
noise = utils.scale_samples(
|
||||
noise,
|
||||
adjusted_shape[-1] * order_spatial_compensation,
|
||||
adjusted_shape[-2] * order_spatial_compensation,
|
||||
mode=upscale_mode,
|
||||
)
|
||||
if self.scatternet_order == 0:
|
||||
return self.fix_output_frames(noise)
|
||||
self.scatternet = self.scatternet.to(device=self.device, dtype=self.dtype)
|
||||
if self.per_channel_scatternet:
|
||||
# To C, B, 1, H, W
|
||||
noise = torch.stack(
|
||||
tuple(
|
||||
self.scatternet(noise[:, chan : chan + 1])
|
||||
for chan in range(self.channels)
|
||||
),
|
||||
dim=0,
|
||||
)
|
||||
else:
|
||||
# To 1, B, C, H, W
|
||||
noise = self.scatternet(noise)[None]
|
||||
base_channels = 1 if self.per_channel_scatternet else self.channels
|
||||
if output_mode == "flat":
|
||||
noise = noise.reshape(noise.shape[0], self.batch, -1)
|
||||
initial_size = math.prod(
|
||||
self.shape[(2 if self.per_channel_scatternet else 1) :],
|
||||
)
|
||||
elif adjusted:
|
||||
initial_size = base_channels
|
||||
else:
|
||||
initial_size = base_channels * ((2**order) ** 2)
|
||||
increment = 1 if output_mode == "flat" else base_channels
|
||||
out_size = noise.shape[2]
|
||||
offset_size = (out_size - initial_size) / increment
|
||||
output_offset = self.output_offset
|
||||
if output_offset == 0 or abs(output_offset) >= 1:
|
||||
output_offset = int(output_offset)
|
||||
if output_offset < 0:
|
||||
output_offset = (offset_size + 1) + output_offset
|
||||
else:
|
||||
if output_offset < 0:
|
||||
output_offset += 1.0
|
||||
output_offset = round(offset_size * output_offset)
|
||||
base_idx = int(output_offset * increment)
|
||||
# print(
|
||||
# f"\nSCAT: shape={noise.shape}, adj_shape={adjusted_shape}, offset={output_offset}, initial_size={initial_size}, out_size={out_size}, offset_size={offset_size}, incr={increment}, base_idx={base_idx}",
|
||||
# )
|
||||
noise = noise[:, :, base_idx : base_idx + initial_size]
|
||||
# print(f"\nSCAT2: {noise.shape}")
|
||||
noise = (
|
||||
noise.squeeze(2).movedim(0, 1) if self.per_channel_scatternet else noise[0]
|
||||
)
|
||||
# print(f"\nSCAT3: {noise.shape}")
|
||||
if output_mode == "channels":
|
||||
noise = noise[..., : self.height, : self.width]
|
||||
# print(
|
||||
# f"\nSCAT4: {noise.shape} -> {adjusted_shape} -- numel: {noise.numel()}, adjnumel={math.prod(adjusted_shape)}",
|
||||
# )
|
||||
return noise.reshape(adjusted_shape).contiguous()
|
||||
@@ -0,0 +1,698 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import operator
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
from comfy.k_diffusion import sampling
|
||||
from torch import FloatTensor, Generator, Tensor
|
||||
from torch.distributions import Laplace, StudentT
|
||||
|
||||
from .. import utils
|
||||
from ..utils import safe_pow, tensor_to
|
||||
|
||||
# ruff: noqa: D413, D417, D212, ANN002, ANN003
|
||||
from .base import FramesToChannelsNoiseGenerator, NoiseError, NoiseGenerator
|
||||
|
||||
|
||||
class GaussianNoiseGenerator(NoiseGenerator):
|
||||
name = "gaussian"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {"normalized": False}
|
||||
|
||||
def generate(self, *_args):
|
||||
return self.rand_like()
|
||||
|
||||
|
||||
class BrownianNoiseGenerator(NoiseGenerator):
|
||||
name = "brownian"
|
||||
|
||||
def __init__(self, x, *args, **kwargs):
|
||||
super().__init__(x, *args, **kwargs)
|
||||
seed = self.options.get("seed")
|
||||
sigma_min = self.options.get("sigma_min")
|
||||
sigma_max = self.options.get("sigma_max")
|
||||
if sigma_min is None or sigma_max is None:
|
||||
raise ValueError("Brownian noise requires sigma_min and sigma_max")
|
||||
self.brownian_tree_ns = sampling.BrownianTreeNoiseSampler(
|
||||
x,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
seed=seed,
|
||||
cpu=self.cpu,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {"normalized": False}
|
||||
|
||||
def generate(self, *args):
|
||||
return self.brownian_tree_ns(*args)
|
||||
|
||||
|
||||
class PerlinOldNoiseGenerator(FramesToChannelsNoiseGenerator):
|
||||
name = "perlin_old"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"div_fac": 2.0,
|
||||
"iterations": 2,
|
||||
"blend_mode": "lerp",
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def get_positions(block_shape: tuple[int, int]) -> Tensor:
|
||||
"""
|
||||
Generate position tensor.
|
||||
|
||||
Arguments:
|
||||
block_shape -- (height, width) of position tensor
|
||||
|
||||
Returns:
|
||||
position vector shaped (1, height, width, 1, 1, 2)
|
||||
"""
|
||||
bh, bw = block_shape
|
||||
return torch.stack(
|
||||
torch.meshgrid(
|
||||
[(torch.arange(b) + 0.5) / b for b in (bw, bh)],
|
||||
indexing="xy",
|
||||
),
|
||||
-1,
|
||||
).view(1, bh, bw, 1, 1, 2)
|
||||
|
||||
@staticmethod
|
||||
def unfold_grid(vectors: Tensor) -> Tensor:
|
||||
"""
|
||||
Unfold vector grid to batched vectors.
|
||||
|
||||
Arguments:
|
||||
vectors -- grid vectors
|
||||
|
||||
Returns:
|
||||
batched grid vectors
|
||||
"""
|
||||
batch_size, _channels, gpy, gpx = vectors.shape
|
||||
return (
|
||||
torch.nn.functional.unfold(vectors, (2, 2))
|
||||
.view(batch_size, 2, 4, -1)
|
||||
.permute(0, 2, 3, 1)
|
||||
.view(batch_size, 4, gpy - 1, gpx - 1, 2)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def smooth_step(t: Tensor) -> Tensor:
|
||||
"""
|
||||
Smooth step function [0, 1] -> [0, 1].
|
||||
|
||||
Arguments:
|
||||
t -- input values (any shape)
|
||||
|
||||
Returns:
|
||||
output values (same shape as input values)
|
||||
"""
|
||||
return t * t * (3.0 - 2.0 * t)
|
||||
|
||||
@classmethod
|
||||
def perlin_noise_tensor(
|
||||
cls,
|
||||
vectors: Tensor,
|
||||
positions: Tensor,
|
||||
step: Callable | None = None,
|
||||
blend=torch.lerp,
|
||||
) -> Tensor:
|
||||
"""
|
||||
Generate perlin noise from batched vectors and positions.
|
||||
|
||||
Arguments:
|
||||
vectors -- batched grid vectors shaped (batch_size, 4, grid_height, grid_width, 2)
|
||||
positions -- batched grid positions shaped (batch_size or 1, block_height, block_width, grid_height or 1, grid_width or 1, 2)
|
||||
|
||||
Keyword Arguments:
|
||||
step -- smooth step function [0, 1] -> [0, 1] (default: `smooth_step`)
|
||||
|
||||
Raises:
|
||||
NoiseError: if position and vector shapes do not match
|
||||
|
||||
Returns:
|
||||
(batch_size, block_height * grid_height, block_width * grid_width)
|
||||
"""
|
||||
if step is None:
|
||||
step = cls.smooth_step
|
||||
|
||||
batch_size = vectors.shape[0]
|
||||
# grid height, grid width
|
||||
gh, gw = vectors.shape[2:4]
|
||||
# block height, block width
|
||||
bh, bw = positions.shape[1:3]
|
||||
|
||||
for i in range(2):
|
||||
if positions.shape[i + 3] not in {1, vectors.shape[i + 2]}:
|
||||
msg = f"Blocks shapes do not match: vectors ({vectors.shape[1]}, {vectors.shape[2]}), positions {gh}, {gw})"
|
||||
raise NoiseError(msg)
|
||||
|
||||
if positions.shape[0] not in {1, batch_size}:
|
||||
msg = f"Batch sizes do not match: vectors ({vectors.shape[0]}), positions ({positions.shape[0]})"
|
||||
raise NoiseError(msg)
|
||||
|
||||
vectors = vectors.view(batch_size, 4, 1, gh * gw, 2)
|
||||
positions = positions.view(positions.shape[0], bh * bw, -1, 2)
|
||||
|
||||
step_x = step(positions[..., 0])
|
||||
step_y = step(positions[..., 1])
|
||||
|
||||
row0 = blend(
|
||||
(vectors[:, 0] * positions).sum(dim=-1),
|
||||
(vectors[:, 1] * (positions - positions.new_tensor((1, 0)))).sum(dim=-1),
|
||||
step_x,
|
||||
)
|
||||
row1 = blend(
|
||||
(vectors[:, 2] * (positions - positions.new_tensor((0, 1)))).sum(dim=-1),
|
||||
(vectors[:, 3] * (positions - positions.new_tensor((1, 1)))).sum(dim=-1),
|
||||
step_x,
|
||||
)
|
||||
noise = blend(row0, row1, step_y)
|
||||
return (
|
||||
noise.view(
|
||||
batch_size,
|
||||
bh,
|
||||
bw,
|
||||
gh,
|
||||
gw,
|
||||
)
|
||||
.permute(0, 3, 1, 4, 2)
|
||||
.reshape(batch_size, gh * bh, gw * bw)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def perlin_noise(
|
||||
cls,
|
||||
grid_shape: tuple[int, int],
|
||||
out_shape: tuple[int, int],
|
||||
batch_size: int = 1,
|
||||
blend=torch.lerp,
|
||||
generator: Generator | None = None,
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> Tensor:
|
||||
"""
|
||||
Generate perlin noise with given shape. `*args` and `**kwargs` are forwarded to `Tensor` creation.
|
||||
|
||||
Arguments:
|
||||
grid_shape -- Shape of grid (height, width).
|
||||
out_shape -- Shape of output noise image (height, width).
|
||||
|
||||
Keyword Arguments:
|
||||
batch_size -- (default: {1})
|
||||
generator -- random generator used for grid vectors (default: {None})
|
||||
|
||||
Raises:
|
||||
NoiseError: if grid and out shapes do not match
|
||||
|
||||
Returns:
|
||||
Noise image shaped (batch_size, height, width)
|
||||
"""
|
||||
# grid height and width
|
||||
gh, gw = grid_shape
|
||||
# output height and width
|
||||
oh, ow = out_shape
|
||||
# block height and width
|
||||
bh, bw = oh // gh, ow // gw
|
||||
|
||||
if oh != bh * gh:
|
||||
msg = f"Output height {oh} must be divisible by grid height {gh}"
|
||||
raise NoiseError(msg)
|
||||
if ow != bw * gw != 0:
|
||||
msg = f"Output width {ow} must be divisible by grid width {gw}"
|
||||
raise NoiseError(msg)
|
||||
|
||||
angle = torch.empty(
|
||||
[batch_size] + [s + 1 for s in grid_shape],
|
||||
*args,
|
||||
**kwargs,
|
||||
).uniform_(to=2.0 * math.pi, generator=generator)
|
||||
# random vectors on grid points
|
||||
vectors = cls.unfold_grid(
|
||||
torch.stack((torch.cos(angle), torch.sin(angle)), dim=1),
|
||||
)
|
||||
# positions inside grid cells [0, 1)
|
||||
positions = tensor_to(cls.get_positions((bh, bw)), vectors)
|
||||
return cls.perlin_noise_tensor(vectors, positions, blend=blend).squeeze(0)
|
||||
|
||||
def generate(self, *_args):
|
||||
blend = utils.BLENDING_MODES[self.blend_mode]
|
||||
noise = self.rand_like(fun=torch.rand).div_(self.div_fac)
|
||||
|
||||
channels, height, width = noise.shape[1:]
|
||||
for _ in range(self.iterations):
|
||||
noise += self.perlin_noise(
|
||||
(height, self.width),
|
||||
(height, width),
|
||||
batch_size=channels,
|
||||
blend=blend,
|
||||
dtype=noise.dtype,
|
||||
layout=noise.layout,
|
||||
device=noise.device,
|
||||
)
|
||||
return self.fix_output_frames(noise)
|
||||
|
||||
|
||||
class UniformNoiseGenerator(NoiseGenerator):
|
||||
name = "uniform"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"normalized": False,
|
||||
"sub_fac": 0.5,
|
||||
"mul_fac": 3.46,
|
||||
"mean_fac": 0.0,
|
||||
}
|
||||
|
||||
def generate(self, *_args):
|
||||
return (
|
||||
self.rand_like(fun=torch.rand)
|
||||
.sub_(self.sub_fac)
|
||||
.mul_(self.mul_fac)
|
||||
.add_(self.mean_fac)
|
||||
)
|
||||
|
||||
|
||||
class HighresPyramidNoiseGenerator(FramesToChannelsNoiseGenerator):
|
||||
name = "highres_pyramid"
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
if self.noise_generator is None:
|
||||
self.noise_generator = UniformNoiseGenerator(
|
||||
*args,
|
||||
**(kwargs | {"normalized": self.normalize_noise}),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"normalized": True,
|
||||
"discount": 0.7,
|
||||
"upscale_mode": "bilinear",
|
||||
"iterations": 4,
|
||||
"noise_generator": None,
|
||||
"normalize_noise": False,
|
||||
}
|
||||
|
||||
def generate(self, s, sn):
|
||||
adjusted_shape = self.get_adjusted_shape()
|
||||
b, c, h, w = adjusted_shape
|
||||
orig_w, orig_h = w, h
|
||||
noise = self.noise_generator(s, sn).reshape(*adjusted_shape)
|
||||
rs = (
|
||||
torch.rand(
|
||||
self.iterations,
|
||||
dtype=torch.float32,
|
||||
generator=self.generator,
|
||||
).cpu()
|
||||
* 2
|
||||
+ 2
|
||||
)
|
||||
for i in range(self.iterations):
|
||||
r = rs[i].item()
|
||||
h, w = min(orig_h * 15, int(h * (r**i))), min(orig_w * 15, int(w * (r**i)))
|
||||
noise += utils.scale_samples(
|
||||
tensor_to(torch.randn(b, c, h, w, generator=self.generator), noise),
|
||||
orig_w,
|
||||
orig_h,
|
||||
mode=self.upscale_mode,
|
||||
).mul_(self.discount**i)
|
||||
if h >= orig_h * 15 or w >= orig_w * 15:
|
||||
break # Lowest resolution is 1x1
|
||||
return self.fix_output_frames(noise)
|
||||
|
||||
|
||||
class PyramidOldNoiseGenerator(FramesToChannelsNoiseGenerator):
|
||||
name = "pyramid_old"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"discount": 0.8,
|
||||
"iterations": 5,
|
||||
"upscale_mode": "nearest-exact",
|
||||
"normalized": False,
|
||||
}
|
||||
|
||||
def generate(self, *_args):
|
||||
adjusted_shape = self.get_adjusted_shape()
|
||||
b, c, h, w = adjusted_shape
|
||||
orig_h, orig_w = h, w
|
||||
noise = torch.zeros(
|
||||
size=adjusted_shape,
|
||||
dtype=self.dtype,
|
||||
layout=self.layout,
|
||||
device=self.gen_device,
|
||||
)
|
||||
r = 1
|
||||
for i in range(self.iterations):
|
||||
r *= 2
|
||||
noise += utils.scale_samples(
|
||||
torch.normal(
|
||||
mean=0,
|
||||
std=0.5**i,
|
||||
size=(b, c, h * r, w * r),
|
||||
dtype=noise.dtype,
|
||||
layout=noise.layout,
|
||||
generator=self.generator,
|
||||
device=noise.device,
|
||||
),
|
||||
orig_w,
|
||||
orig_h,
|
||||
mode=self.upscale_mode,
|
||||
).mul_(self.discount**i)
|
||||
return self.fix_output_frames(noise)
|
||||
|
||||
|
||||
class PyramidNoiseGenerator(FramesToChannelsNoiseGenerator):
|
||||
name = "pyramid"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"discount": 0.7,
|
||||
"upscale_mode": "bilinear",
|
||||
"iterations": 10,
|
||||
"iteration_offset": 0,
|
||||
"iteration_step": 1,
|
||||
"reverse_scale": False,
|
||||
"reverse_size_h": False,
|
||||
"reverse_size_w": False,
|
||||
"base_h": 2.0,
|
||||
"multiplier_h": 2.0,
|
||||
"base_w": 2.0,
|
||||
"multiplier_w": 2.0,
|
||||
"legacy_r": False,
|
||||
"size_min": 1,
|
||||
"size_max_pct": 2.0,
|
||||
"include_size_limit": False,
|
||||
"high_res_mode": False,
|
||||
}
|
||||
|
||||
# Original implementatino modified from https://wandb.ai/johnowhitaker/multires_noise/reports/Multi-Resolution-Noise-for-Diffusion-Model-Training--VmlldzozNjYyOTU2
|
||||
def generate(self, *_args):
|
||||
noise = self.rand_like()
|
||||
b, c, h, w = noise.shape
|
||||
orig_w, orig_h = w, h
|
||||
size_min = max(1, self.size_min)
|
||||
max_h = max(1, int(orig_h * self.size_max_pct))
|
||||
max_w = max(1, int(orig_w * self.size_max_pct))
|
||||
eps = 1e-02
|
||||
op = operator.mul if self.high_res_mode else operator.truediv
|
||||
|
||||
if self.legacy_r:
|
||||
|
||||
def get_r(_i: int) -> float:
|
||||
return torch.rand(1, generator=self.generator).cpu().item()
|
||||
else:
|
||||
rs = torch.rand(self.iterations, generator=self.generator).cpu().tolist()
|
||||
|
||||
def get_r(i: int) -> float:
|
||||
return rs[i]
|
||||
|
||||
for i in range(
|
||||
self.iteration_offset,
|
||||
self.iterations + self.iteration_offset,
|
||||
self.iteration_step,
|
||||
):
|
||||
rev_i = self.iterations - i - 1
|
||||
r = get_r(i)
|
||||
rh = r * self.multiplier_h + self.base_h
|
||||
rw = r * self.multiplier_w + self.base_w
|
||||
ih = rev_i if self.reverse_size_h else i
|
||||
iw = rev_i if self.reverse_size_w else i
|
||||
h = max(1, min(max_h, int(op(h, max(eps, rh**ih)))))
|
||||
w = max(1, min(max_w, int(op(w, max(eps, rw**iw)))))
|
||||
size_limit = h <= size_min or w <= size_min or h >= max_h or w >= max_w
|
||||
if not self.include_size_limit and size_limit:
|
||||
break
|
||||
scale = self.discount ** (rev_i if self.reverse_scale else i)
|
||||
if scale == 0:
|
||||
continue
|
||||
noise += utils.scale_samples(
|
||||
torch.randn(
|
||||
b,
|
||||
c,
|
||||
h,
|
||||
w,
|
||||
device=noise.device,
|
||||
layout=noise.layout,
|
||||
dtype=noise.dtype,
|
||||
),
|
||||
orig_w,
|
||||
orig_h,
|
||||
mode=self.upscale_mode,
|
||||
).mul_(scale)
|
||||
if size_limit:
|
||||
break
|
||||
return self.fix_output_frames(noise)
|
||||
|
||||
|
||||
class StudentTNoiseGenerator(NoiseGenerator):
|
||||
name = "studentt"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"loc": 0,
|
||||
"scale": 0.2,
|
||||
"df": 1,
|
||||
"quantile_fac": 0.75,
|
||||
"pow_fac": 0.5,
|
||||
"nq_fac": 1.0,
|
||||
"normalized": False,
|
||||
}
|
||||
|
||||
def generate(self, *_args):
|
||||
noise = StudentT(loc=self.loc, scale=self.scale, df=self.df).rsample(self.shape)
|
||||
nq = torch.quantile(
|
||||
noise.flatten(start_dim=1).abs(),
|
||||
self.quantile_fac,
|
||||
dim=-1,
|
||||
)
|
||||
nq_shape = tuple(nq.shape) + (1,) * (noise.ndim - nq.ndim)
|
||||
nq = nq.mul_(self.nq_fac).reshape(*nq_shape)
|
||||
noise = noise.clamp_(-nq, nq)
|
||||
return noise.abs().pow_(self.pow_fac).copysign_(noise)
|
||||
|
||||
|
||||
class GreenTestNoiseGenerator(FramesToChannelsNoiseGenerator):
|
||||
name = "green_test"
|
||||
MIN_DIMS = 4
|
||||
MAX_DIMS = 5
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"scale_fac": 1.0,
|
||||
"x_pow": 2.0,
|
||||
"y_pow": 2.0,
|
||||
"x_multiplier": 1.0,
|
||||
"y_multiplier": 1.0,
|
||||
"power_base": 1.0,
|
||||
"inv_power": 0.5,
|
||||
"restore_sign_power": False,
|
||||
"restore_sign_x": False,
|
||||
"restore_sign_y": False,
|
||||
}
|
||||
|
||||
def generate(self, *_args):
|
||||
noise = self.rand_like()
|
||||
scale = self.scale_fac / max(1, self.width * self.height)
|
||||
fy, fx = (
|
||||
torch.fft.fftfreq(sz, device=noise.device, dtype=noise.dtype)
|
||||
for sz in (self.height, self.width)
|
||||
)
|
||||
fx = safe_pow(fx, self.x_pow, restore_sign=self.restore_sign_x, in_place=True)
|
||||
fy = safe_pow(fy, self.y_pow, restore_sign=self.restore_sign_y, in_place=True)
|
||||
if self.x_multiplier != 1:
|
||||
fx *= self.x_multiplier
|
||||
if self.y_multiplier != 1:
|
||||
fy *= self.y_multiplier
|
||||
power = fy[:, None] + fx
|
||||
inv_power = self.inv_power * self.inv_power
|
||||
power = safe_pow(
|
||||
power,
|
||||
inv_power,
|
||||
restore_sign=self.restore_sign_power,
|
||||
in_place=True,
|
||||
)
|
||||
coord_0 = self.power_base**self.inv_power
|
||||
if coord_0 == 0 or not math.isfinite(coord_0):
|
||||
coord_0 = 1.0
|
||||
power = power.masked_fill_((power == 0) | (~power.isfinite()), coord_0)
|
||||
power[0, 0] = coord_0
|
||||
noise *= scale
|
||||
noise = torch.fft.ifft2(torch.fft.fft2(noise).div_(power))
|
||||
return self.fix_output_frames(noise.real)
|
||||
|
||||
|
||||
class PinkOldNoiseGenerator(NoiseGenerator):
|
||||
name = "pink_old"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {"alpha": 2.0, "k": 1.0, "freq": 1.0}
|
||||
|
||||
# Completely wrong implementation here.
|
||||
def generate(self, *_args):
|
||||
spectral_density = self.k / self.freq**self.alpha
|
||||
return self.rand_like() * spectral_density
|
||||
|
||||
|
||||
def frequency_scaled_noise(
|
||||
x: torch.Tensor,
|
||||
*,
|
||||
x_is_noise: bool = False,
|
||||
base_power: float = 0.5,
|
||||
alpha: float,
|
||||
) -> torch.Tensor:
|
||||
h, w = x.shape[-2:]
|
||||
fh = torch.fft.fftfreq(h, device=x.device).unsqueeze(-1)
|
||||
fw = torch.fft.fftfreq(w, device=x.device).unsqueeze(0)
|
||||
p = (fh**2 + fw**2).pow_(base_power * alpha)
|
||||
p[0, 0] = 1.0**alpha
|
||||
noise = x if x_is_noise else torch.randn_like(x)
|
||||
noise_fft = torch.fft.fftn(noise, dim=(-2, -1))
|
||||
p = p.to(noise_fft.dtype).expand(*((1,) * (x.ndim - 2)), h, w)
|
||||
noise_fft /= p
|
||||
noise_fft[..., 0, 0] = 0.0
|
||||
noise = torch.fft.ifftn(noise_fft, dim=(-2, -1)).real.to(x.dtype)
|
||||
noise /= noise.std(dim=tuple(range(1, x.ndim)), keepdim=True).clamp_min_(1e-06)
|
||||
return noise
|
||||
|
||||
|
||||
class OneFNoiseGenerator(FramesToChannelsNoiseGenerator):
|
||||
name = "onef"
|
||||
MIN_DIMS = 4
|
||||
MAX_DIMS = 5
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"alpha": 2.0,
|
||||
"k": 1.0,
|
||||
"hfac": 1.0,
|
||||
"wfac": 1.0,
|
||||
"base_power": 1.0,
|
||||
"use_sqrt": True,
|
||||
# None or or float, alternative to use_sqrt with custom power.
|
||||
"power": None,
|
||||
"x_pow": 2.0,
|
||||
"y_pow": 2.0,
|
||||
}
|
||||
|
||||
# Original implementation referenced from: https://github.com/WASasquatch/PowerNoiseSuite
|
||||
def generate(self, *_args):
|
||||
noise = self.rand_like()
|
||||
|
||||
freq_x, freq_y = (
|
||||
torch.fft.fftfreq(sz, fac, device=noise.device, dtype=noise.dtype)
|
||||
for sz, fac in ((self.height, self.hfac), (self.width, self.wfac))
|
||||
)
|
||||
freq_x **= self.x_pow
|
||||
freq_y **= self.y_pow
|
||||
fx, fy = torch.meshgrid(freq_x, freq_y, indexing="ij")
|
||||
power = fx + fy
|
||||
power **= self.alpha / -2.0
|
||||
if self.k not in {0, 1}:
|
||||
power *= 1 / self.k
|
||||
noise_fft = torch.fft.fftn(noise)
|
||||
user_power = 0.5 if self.use_sqrt else self.power
|
||||
if isinstance(user_power, float):
|
||||
power **= user_power
|
||||
coord_0 = self.base_power
|
||||
if coord_0 == 0 or not math.isfinite(coord_0):
|
||||
coord_0 = 1.0
|
||||
power = power.masked_fill_((power == 0) | (~power.isfinite()), coord_0)
|
||||
power = (
|
||||
power.to(dtype=noise_fft.dtype)
|
||||
.unsqueeze(0)
|
||||
.expand(self.batch, 1, self.height, self.width)
|
||||
)
|
||||
noise_fft /= power
|
||||
noise = torch.fft.ifftn(noise_fft).real
|
||||
return self.fix_output_frames(noise)
|
||||
|
||||
|
||||
class PowerLawNoiseGenerator(NoiseGenerator):
|
||||
name = "powerlaw"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"alpha": 2.0,
|
||||
"div_max_dims": None,
|
||||
"use_sign": False,
|
||||
"use_div_max_abs": True,
|
||||
}
|
||||
|
||||
# Referenced from: https://github.com/WASasquatch/PowerNoiseSuite
|
||||
def generate(self, *_args):
|
||||
noise = self.rand_like()
|
||||
|
||||
modulation = torch.abs(noise) ** self.alpha
|
||||
noise = (torch.sign(noise) if self.use_sign else noise).mul_(modulation)
|
||||
if self.div_max_dims is not None:
|
||||
noise /= torch.amax(
|
||||
torch.abs(noise) if self.use_div_max_abs else noise,
|
||||
keepdim=True,
|
||||
dim=self.div_max_dims,
|
||||
)
|
||||
return noise
|
||||
|
||||
|
||||
class LaplacianNoiseGenerator(NoiseGenerator):
|
||||
name = "laplacian"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {"loc": 0, "scale": 1.0, "div_fac": 4.0}
|
||||
|
||||
def generate(self, *_args):
|
||||
noise = self.rand_like().div_(self.div_fac)
|
||||
noise += tensor_to(
|
||||
Laplace(loc=self.loc, scale=self.scale).rsample(self.shape),
|
||||
noise.device,
|
||||
)
|
||||
return noise
|
||||
|
||||
|
||||
class PowerOldNoiseGenerator(NoiseGenerator):
|
||||
name = "power_old"
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {"alpha": 2, "k": 1, "normalized": False}
|
||||
|
||||
def generate(self, *_args):
|
||||
tensor = self.rand_like()
|
||||
fft = torch.fft.fft2(tensor)
|
||||
freq = torch.arange(
|
||||
1,
|
||||
len(fft) + 1,
|
||||
dtype=tensor.dtype,
|
||||
layout=tensor.layout,
|
||||
device=tensor.device,
|
||||
).reshape(
|
||||
(len(fft),) + (1,) * (tensor.dim() - 1),
|
||||
)
|
||||
spectral_density = self.k / freq**self.alpha
|
||||
noise = torch.rand(
|
||||
tensor.shape,
|
||||
device=tensor.device,
|
||||
layout=tensor.layout,
|
||||
dtype=tensor.dtype,
|
||||
).mul_(spectral_density)
|
||||
mean = torch.mean(noise, dim=(-2, -1), keepdim=True)
|
||||
std = torch.std(noise, dim=(-2, -1), keepdim=True)
|
||||
return noise.sub_(mean).div_(std)
|
||||
@@ -0,0 +1,597 @@
|
||||
# ruff: noqa: ANN002, ANN003
|
||||
from __future__ import annotations
|
||||
|
||||
import itertools
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from .base import NoiseGenerator
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
|
||||
class SimulationNoiseGenerator(NoiseGenerator):
|
||||
name = "simulation"
|
||||
MIN_DIMS = 4
|
||||
MAX_DIMS = 4
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls, *, no_super: bool = False):
|
||||
result = {
|
||||
# multi_octave, power_law, band_pass
|
||||
"spectral_mode": "multi_octave",
|
||||
# curl, projection, basis
|
||||
"field_mode": "basis",
|
||||
"depth_mode": "reset",
|
||||
"channel_mode": "stacked",
|
||||
"band_shape": "log_gaussian",
|
||||
"dims": (),
|
||||
"base_k": 0.0,
|
||||
"power_law_beta": 1.0,
|
||||
"depth": 64,
|
||||
"initial_depth": 0,
|
||||
"max_depth": -1,
|
||||
# reset, wrap, bounce
|
||||
"octaves": 5,
|
||||
"lacunarity": 2.0,
|
||||
"gain": 0.5,
|
||||
# log_gaussian, raised_cosine
|
||||
"log_gaussian_sigma": 0.3,
|
||||
# (float, float, float)
|
||||
"band_pass_low": 0.00001,
|
||||
"band_pass_high": 1.0,
|
||||
"anisotropy": (),
|
||||
"normalized": False,
|
||||
"noise_sampler_factory_h": None,
|
||||
"noise_sampler_factory_w": None,
|
||||
"noise_sampler_factory_z": None,
|
||||
}
|
||||
return result if no_super else super().ng_params() | result
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.noise_chunk = None
|
||||
cm = self.channel_mode
|
||||
self.depth_increment = 1
|
||||
if cm in {"over_depth", "over_depth_alt"}:
|
||||
self.depth_increment = math.ceil(self.channels / 3)
|
||||
elif cm.startswith("over_depth_"):
|
||||
self.depth_increment = self.channels
|
||||
else:
|
||||
self.depth_increment = 1
|
||||
if self.initial_depth < 0:
|
||||
self.initial_depth = self.depth + self.initial_depth
|
||||
if self.initial_depth < 0:
|
||||
raise ValueError("Initial depth out of range")
|
||||
self.initial_depth = min(self.depth - 1, self.initial_depth)
|
||||
if self.max_depth < 0:
|
||||
self.max_depth = self.depth + self.max_depth
|
||||
if self.max_depth < 0:
|
||||
raise ValueError("Max depth out of range")
|
||||
self.max_depth = min(self.depth - 1, self.max_depth)
|
||||
self.current_depth = self.initial_depth
|
||||
self.direction = 1
|
||||
self.cdtype = (
|
||||
(torch.complex128 if self.dtype == torch.float64 else torch.complex64)
|
||||
if not self.dtype.is_complex
|
||||
else self.dtype
|
||||
)
|
||||
self.eff_batch = (
|
||||
self.batch
|
||||
if cm not in {"stacked", "flat"}
|
||||
else self.batch * math.ceil(self.channels / 3)
|
||||
)
|
||||
ns_shape = torch.Size(
|
||||
(
|
||||
self.eff_batch,
|
||||
self.depth * self.depth_increment,
|
||||
self.height,
|
||||
self.width,
|
||||
)
|
||||
)
|
||||
|
||||
def gaussian_noise_sampler(*_args: Any) -> torch.Tensor:
|
||||
return torch.randn(ns_shape, dtype=self.cdtype, device=self.gen_device).to(
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
self.noise_samplers = tuple(
|
||||
factory.make_noise_sampler(
|
||||
torch.zeros(ns_shape, device=self.gen_device, dtype=self.cdtype),
|
||||
cpu=self.cpu,
|
||||
normalized=False,
|
||||
)
|
||||
if factory is not None
|
||||
else gaussian_noise_sampler
|
||||
for factory in (
|
||||
self.noise_sampler_factory_z,
|
||||
self.noise_sampler_factory_h,
|
||||
self.noise_sampler_factory_w,
|
||||
)
|
||||
)
|
||||
|
||||
def _k_grids(self, *, shape: tuple, dims: tuple = (-3, -2, -1)) -> tuple:
|
||||
"""Creates k-space grids."""
|
||||
return torch.meshgrid(
|
||||
*(
|
||||
torch.fft.fftfreq(
|
||||
shape[dim],
|
||||
d=1.0,
|
||||
device=self.device,
|
||||
dtype=self.dtype if not self.dtype.is_complex else torch.float64,
|
||||
).to(dtype=self.dtype)
|
||||
for dim in dims
|
||||
),
|
||||
indexing="ij",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _radial_k(*ks: torch.Tensor) -> torch.Tensor:
|
||||
"""Calculates the radial distance in k-space."""
|
||||
return sum(kt**2 for kt in ks).sqrt_()
|
||||
|
||||
@staticmethod
|
||||
def _raised_cosine_band(
|
||||
k: torch.Tensor,
|
||||
k_lo: float,
|
||||
k_hi: float,
|
||||
) -> torch.Tensor:
|
||||
"""A raised cosine spectral band filter."""
|
||||
kc = 0.5 * (k_lo + k_hi)
|
||||
hw = 0.5 * (k_hi - k_lo) + 1e-12
|
||||
t = (k - kc) / hw
|
||||
return torch.where(
|
||||
t.abs() <= 1.0,
|
||||
0.5 * (1.0 + torch.cos(math.pi * t)),
|
||||
torch.zeros_like(k),
|
||||
)
|
||||
|
||||
def _log_gaussian_band(
|
||||
self,
|
||||
k: torch.Tensor,
|
||||
k_lo: float,
|
||||
k_hi: float,
|
||||
) -> torch.Tensor:
|
||||
"""A log-Gaussian spectral band filter."""
|
||||
k_center = math.sqrt(k_lo * k_hi)
|
||||
log_k = torch.log(torch.clamp(k, min=1e-12))
|
||||
log_center = math.log(k_center)
|
||||
return torch.exp(-0.5 * ((log_k - log_center) / self.log_gaussian_sigma) ** 2)
|
||||
|
||||
def _handle_band_shape(
|
||||
self,
|
||||
k_rad: torch.Tensor,
|
||||
k_low: float,
|
||||
k_high: float,
|
||||
) -> torch.Tensor:
|
||||
if self.band_shape == "raised_cosine":
|
||||
return self._raised_cosine_band(k_rad, k_low, k_high)
|
||||
if self.band_shape == "log_gaussian":
|
||||
return self._log_gaussian_band(k_rad, k_low, k_high)
|
||||
errstr = f"Bad band shape mode {self.band_shape}"
|
||||
raise ValueError(errstr)
|
||||
|
||||
def _make_wk(
|
||||
self,
|
||||
k_rad: torch.Tensor,
|
||||
sizes: tuple,
|
||||
*,
|
||||
eps: float = 1e-09,
|
||||
) -> torch.Tensor:
|
||||
def wk_out(wk: torch.Tensor) -> torch.Tensor:
|
||||
wk[k_rad == 0] = 0.0
|
||||
return wk
|
||||
|
||||
if self.spectral_mode == "power_law":
|
||||
return wk_out((k_rad + eps).pow_(-self.power_law_beta))
|
||||
if self.spectral_mode == "band_pass":
|
||||
if self.band_pass_low >= self.band_pass_high:
|
||||
raise ValueError(
|
||||
"band_pass_high must be greater than band_pass_low in band_pass spectral mode.",
|
||||
)
|
||||
return wk_out(
|
||||
self._handle_band_shape(k_rad, self.band_pass_low, self.band_pass_high),
|
||||
)
|
||||
if self.octaves == 0:
|
||||
# Ones where k_rad is non-zero, otherwise zero.
|
||||
return (k_rad != 0).to(k_rad)
|
||||
base_k = 2 * math.pi / max(1, min(sizes)) if self.base_k == 0 else self.base_k
|
||||
wk = torch.zeros_like(k_rad)
|
||||
for o in range(self.octaves):
|
||||
k_lo = base_k * (self.lacunarity**o)
|
||||
k_hi = base_k * (self.lacunarity ** (o + 1))
|
||||
band = self._handle_band_shape(k_rad, k_lo, k_hi)
|
||||
wk += (self.gain**o) * band
|
||||
|
||||
return wk_out(wk)
|
||||
|
||||
def _handle_field_projection(
|
||||
self,
|
||||
*,
|
||||
k_grids_orig: tuple,
|
||||
wk: torch.Tensor,
|
||||
ns_args: tuple | list,
|
||||
**_kwargs,
|
||||
):
|
||||
n_dims = len(k_grids_orig)
|
||||
n_samplers = len(self.noise_samplers)
|
||||
f_fs = tuple(
|
||||
self.noise_samplers[ns_idx % n_samplers](*ns_args)
|
||||
.to(
|
||||
device=self.device,
|
||||
)
|
||||
.mul_(wk)
|
||||
for ns_idx in range(n_dims)
|
||||
)
|
||||
|
||||
# --- Perform the Helmholtz projection using the UN SCALED grids ---
|
||||
k_sq_proj = self._radial_k(*k_grids_orig) ** 2
|
||||
k_dot_f = sum(k_p * f_f for k_p, f_f in zip(k_grids_orig, f_fs))
|
||||
inv_k_sq = torch.where(k_sq_proj == 0, 0.0, 1.0 / k_sq_proj)
|
||||
|
||||
k_grid_scale = k_dot_f.mul_(inv_k_sq)
|
||||
return tuple(
|
||||
f_f - k_grid * k_grid_scale for f_f, k_grid in zip(f_fs, k_grids_orig)
|
||||
)
|
||||
|
||||
def _handle_field_curl(
|
||||
self,
|
||||
*,
|
||||
k_rad: torch.Tensor,
|
||||
k_grids_orig: tuple,
|
||||
wk: torch.Tensor,
|
||||
ns_args: tuple | list,
|
||||
**_kwargs,
|
||||
):
|
||||
n_dims = len(k_grids_orig)
|
||||
nd_fixup = int(self.field_mode != "curl_ndim")
|
||||
|
||||
# The potential filter still uses the scaled k_rad for spectral shaping
|
||||
inv_k_rad = torch.where(k_rad == 0, 0.0, 1.0 / k_rad)
|
||||
wk_potential = wk * inv_k_rad
|
||||
|
||||
# The curl operator (i*k) MUST use the original, un-scaled grids
|
||||
i_k_grids = tuple(
|
||||
(1j * k_grid).to(dtype=self.cdtype) for k_grid in k_grids_orig
|
||||
)
|
||||
|
||||
n_samplers = len(self.noise_samplers)
|
||||
g_fs = tuple(
|
||||
self.noise_samplers[ns_idx % n_samplers](*ns_args)
|
||||
.to(
|
||||
device=self.device,
|
||||
)
|
||||
.mul_(wk_potential)
|
||||
for ns_idx in range(n_dims if n_dims != 2 else 1)
|
||||
)
|
||||
# --- Case 1: 2D Curl (Curl of a SCALAR potential) ---
|
||||
# This is the fundamental building block.
|
||||
if n_dims * nd_fixup == 2:
|
||||
# We only need one scalar potential field G.
|
||||
g_f = g_fs[0]
|
||||
ikx, iky = i_k_grids
|
||||
# F = (dG/dy, -dG/dx) -> F_f = (iky*G_f, -ikx*G_f)
|
||||
return (iky * g_f, -ikx * g_f)
|
||||
|
||||
# --- Case 2: 3D Curl (The classic cross-product) ---
|
||||
# This is a special, unique case.
|
||||
if n_dims * nd_fixup == 3:
|
||||
gz_f, gy_f, gx_f = g_fs
|
||||
ikz, iky, ikx = i_k_grids
|
||||
# F_f = i*k x G_f
|
||||
return (
|
||||
ikx * gy_f - iky * gx_f, # z component
|
||||
ikz * gx_f - ikx * gz_f, # y component
|
||||
iky * gz_f - ikz * gy_f, # x component
|
||||
)
|
||||
|
||||
# --- Case 3: N-D Curl (Pragmatic construction) ---
|
||||
# We build the N-D field by summing 2D curls on orthogonal planes.
|
||||
|
||||
# We need N potential fields, but we will use them in pairs.
|
||||
|
||||
f_f_outputs = [torch.zeros_like(g_fs[0]) for _ in range(n_dims)]
|
||||
|
||||
# Iterate over pairs of dimensions (0,1), (2,3), etc.
|
||||
for i in range(n_dims // 2):
|
||||
idx1 = i * 2
|
||||
idx2 = i * 2 + 1
|
||||
|
||||
g1_f = g_fs[idx1]
|
||||
g2_f = g_fs[idx2]
|
||||
ik1 = i_k_grids[idx1]
|
||||
ik2 = i_k_grids[idx2]
|
||||
|
||||
# Perform a 2D-like curl on the (G1, G2) plane
|
||||
# This is a bit abstract, but we are creating rotation in the 1-2 plane.
|
||||
# f_f_outputs[idx1] = ik2 * g1_f - ik1 * g2_f
|
||||
# f_f_outputs[idx2] = ik1 * g2_f - ik2 * g1_f
|
||||
f_f_outputs[idx1] = ik2 * g1_f - ik1 * g2_f
|
||||
f_f_outputs[idx2] = -ik1 * g1_f - ik2 * g2_f
|
||||
|
||||
return tuple(f_f_outputs)
|
||||
|
||||
_handle_field_curl_ndim = _handle_field_curl
|
||||
|
||||
def _handle_field_basis(
|
||||
self,
|
||||
*,
|
||||
k_grids_orig: tuple,
|
||||
wk: torch.Tensor,
|
||||
ns_args: tuple | list,
|
||||
**_kwargs,
|
||||
) -> tuple:
|
||||
n_dims = len(k_grids_orig)
|
||||
nd_fixup = int(self.field_mode != "basis_ndim")
|
||||
k_rad_orig = self._radial_k(*k_grids_orig)
|
||||
|
||||
# Normalize the original k vector
|
||||
k_norm_components = tuple(
|
||||
torch.where(k_rad_orig == 0, 0.0, k / k_rad_orig) for k in k_grids_orig
|
||||
)
|
||||
|
||||
# --- Case 1: 2D (simple and fast) ---
|
||||
if n_dims * nd_fixup == 2:
|
||||
# The basis is a single vector perpendicular to k: u = (-ky, kx)
|
||||
kn_y, kn_x = k_norm_components
|
||||
basis_vectors = [
|
||||
(-kn_x, kn_y),
|
||||
] # A list containing one basis vector (a tuple)
|
||||
|
||||
num_random_fields = 1
|
||||
|
||||
# --- Case 2: 3D (fast cross-product method) ---
|
||||
elif n_dims * nd_fixup == 3:
|
||||
num_random_fields = 2
|
||||
|
||||
kn_z, kn_y, kn_x = k_norm_components
|
||||
ez = torch.tensor([0.0, 0.0, 1.0], device=self.device, dtype=self.dtype)
|
||||
is_parallel = (kn_x.abs() < 1e-6) & (kn_y.abs() < 1e-6)
|
||||
|
||||
ux = torch.where(is_parallel, 0.0, kn_y * ez[2] - kn_z * ez[1])
|
||||
uy = torch.where(is_parallel, -kn_z, kn_z * ez[0] - kn_x * ez[2])
|
||||
uz = torch.where(is_parallel, kn_x, kn_x * ez[1] - kn_y * ez[0])
|
||||
|
||||
u_mag = torch.sqrt(ux**2 + uy**2 + uz**2)
|
||||
inv_u_mag = torch.where(u_mag == 0, 0.0, 1.0 / u_mag)
|
||||
ux, uy, uz = ux * inv_u_mag, uy * inv_u_mag, uz * inv_u_mag
|
||||
|
||||
# u = (uz, uy, ux)
|
||||
u = (ux, uy, uz)
|
||||
|
||||
vx = kn_y * u[2] - kn_z * u[1]
|
||||
vy = kn_z * u[0] - kn_x * u[2]
|
||||
vz = kn_x * u[1] - kn_y * u[0]
|
||||
|
||||
v = (vz, vy, vx)
|
||||
|
||||
basis_vectors = [u, v]
|
||||
|
||||
# --- Case 3: N-D (General Gram-Schmidt process) ---
|
||||
else:
|
||||
num_random_fields = n_dims - 1
|
||||
|
||||
basis_vectors = []
|
||||
# Start with the standard basis vectors (e.g., [1,0,0], [0,1,0], [0,0,1])
|
||||
for i in range(n_dims):
|
||||
# Create a standard basis vector e_i
|
||||
e_i = [torch.zeros_like(k_rad_orig) for _ in range(n_dims)]
|
||||
e_i[i] = torch.ones_like(k_rad_orig)
|
||||
|
||||
# Start with v = e_i and make it orthogonal to k
|
||||
v = list(e_i)
|
||||
dot_k = sum(
|
||||
v_comp * k_comp for v_comp, k_comp in zip(v, k_norm_components)
|
||||
)
|
||||
v = [
|
||||
v_comp - dot_k * k_comp
|
||||
for v_comp, k_comp in zip(v, k_norm_components)
|
||||
]
|
||||
|
||||
# Make it orthogonal to all previously found basis vectors
|
||||
for b in basis_vectors:
|
||||
dot_b = sum(v_comp * b_comp for v_comp, b_comp in zip(v, b))
|
||||
v = [v_comp - dot_b * b_comp for v_comp, b_comp in zip(v, b)]
|
||||
|
||||
# Normalize the new basis vector
|
||||
v_mag = torch.sqrt(sum(comp**2 for comp in v))
|
||||
# Only add the vector if it's not a zero vector
|
||||
if torch.any(v_mag > 1e-6):
|
||||
inv_v_mag = torch.where(v_mag == 0, 0.0, 1.0 / v_mag)
|
||||
v = [comp * inv_v_mag for comp in v]
|
||||
basis_vectors.append(tuple(v))
|
||||
|
||||
if len(basis_vectors) == num_random_fields:
|
||||
break
|
||||
|
||||
# --- Field Construction (works for all cases) ---
|
||||
|
||||
# Generate N-1 independent random complex scalar fields
|
||||
n_samplers = len(self.noise_samplers)
|
||||
random_fields = tuple(
|
||||
self.noise_samplers[ns_idx % n_samplers](*ns_args)
|
||||
.to(device=self.device)
|
||||
.mul_(wk)
|
||||
for ns_idx in range(num_random_fields)
|
||||
)
|
||||
|
||||
# Initialize the final field components to zero
|
||||
f_f_outputs = [torch.zeros_like(random_fields[0]) for _ in range(n_dims)]
|
||||
|
||||
# Project each random field onto its corresponding basis vector and sum them up
|
||||
for i in range(num_random_fields):
|
||||
a_f = random_fields[i]
|
||||
basis_vec = basis_vectors[i]
|
||||
for j in range(n_dims):
|
||||
f_f_outputs[j] += a_f * basis_vec[j]
|
||||
|
||||
return tuple(f_f_outputs)
|
||||
|
||||
_handle_field_basis_ndim = _handle_field_basis
|
||||
|
||||
def calculate_spectral_divergence_3d(
|
||||
self,
|
||||
field: torch.Tensor,
|
||||
*,
|
||||
debug: bool = False,
|
||||
) -> torch.Tensor:
|
||||
if field.ndim != 5:
|
||||
errstr = f"Field must be 5d, got shape {field.shape}"
|
||||
raise ValueError(errstr)
|
||||
C = field.shape[1]
|
||||
if C != 3:
|
||||
errstr = f"Field must have 3 channels, but has {C}"
|
||||
raise ValueError(errstr)
|
||||
cdtype = torch.complex128 if field.dtype == torch.float64 else torch.complex64
|
||||
KX, KY, KZ = (t.to(field) for t in self._k_grids(shape=field.shape))
|
||||
fx_f = torch.fft.fftn(field[:, 0, ...], dim=(-3, -2, -1))
|
||||
fy_f = torch.fft.fftn(field[:, 1, ...], dim=(-3, -2, -1))
|
||||
fz_f = torch.fft.fftn(field[:, 2, ...], dim=(-3, -2, -1))
|
||||
div_f = (
|
||||
(1j * KX.to(cdtype)) * fx_f
|
||||
+ (1j * KY.to(cdtype)) * fy_f
|
||||
+ (1j * KZ.to(cdtype)) * fz_f
|
||||
)
|
||||
result = torch.fft.ifftn(div_f, dim=(-3, -2, -1)).real
|
||||
divergences = result.abs_().mean(dim=tuple(range(1, result.ndim)))
|
||||
if not debug:
|
||||
return divergences
|
||||
prettydivs = ", ".join(
|
||||
f"{dm:.5f}" for dm in divergences.detach().cpu().tolist()
|
||||
)
|
||||
tqdm.write(
|
||||
f"Simulation noise: Input shape: {field.shape}, Mean Absolute Divergences (per batch): {prettydivs}",
|
||||
)
|
||||
return divergences
|
||||
|
||||
def generate_field(
|
||||
self,
|
||||
batch: int,
|
||||
height: int,
|
||||
width: int,
|
||||
*,
|
||||
ns_args: tuple | list,
|
||||
) -> torch.Tensor:
|
||||
depth = self.depth * self.depth_increment
|
||||
eff_shape = torch.Size((batch, 3, depth, height, width))
|
||||
|
||||
# 1. Create the UN SCALED k-grids for the projection operator.
|
||||
k_grids_orig = k_grids = self._k_grids(shape=eff_shape)
|
||||
|
||||
# 2. Create a separate set of k-grids for spectral shaping.
|
||||
# These can be scaled by the anisotropy factors.
|
||||
if self.anisotropy and not all(v in {0, 1} for v in self.anisotropy):
|
||||
n_anisotropy = len(self.anisotropy)
|
||||
anisotropy = tuple(
|
||||
1.0 if idx >= n_anisotropy else self.anisotropy[idx] for idx in range(3)
|
||||
)
|
||||
k_grids = tuple(
|
||||
k_p if a in {None, 0, 1} else k_p / a
|
||||
for k_p, a in itertools.zip_longest(k_grids, anisotropy)
|
||||
)
|
||||
|
||||
# 3. Calculate radial k for the spectral envelope using the SCALED grids.
|
||||
k_rad = self._radial_k(*k_grids)
|
||||
|
||||
# Build the multi-octave spectral envelope (Wk) using the anisotropic k_rad
|
||||
wk = self._make_wk(k_rad, sizes=(depth, height, width))
|
||||
|
||||
field_handler = getattr(self, f"_handle_field_{self.field_mode}", None)
|
||||
if field_handler is None:
|
||||
errstr = f"Bad field mode {self.field_mode}"
|
||||
raise ValueError(errstr)
|
||||
f_f_outputs = field_handler(
|
||||
k_rad=k_rad,
|
||||
k_grids=k_grids,
|
||||
k_grids_orig=k_grids_orig,
|
||||
wk=wk,
|
||||
ns_args=ns_args,
|
||||
)
|
||||
|
||||
# Inverse FFT to transform the field back to the spatial domain
|
||||
fields = tuple(
|
||||
torch.fft.ifftn(f_proj, dim=(-3, -2, -1)).real
|
||||
for f_proj in reversed(f_f_outputs)
|
||||
)
|
||||
field = torch.stack(fields, dim=1)
|
||||
self.calculate_spectral_divergence_3d(field, debug=True)
|
||||
|
||||
rms = torch.sqrt(torch.mean(field**2))
|
||||
if rms > 1e-9:
|
||||
field /= rms
|
||||
|
||||
return field
|
||||
|
||||
def generate(self, *args) -> torch.Tensor:
|
||||
cm = self.channel_mode
|
||||
if self.noise_chunk is None:
|
||||
self.noise_chunk = self.generate_field(
|
||||
self.eff_batch,
|
||||
self.height,
|
||||
self.width,
|
||||
ns_args=args,
|
||||
).to(dtype=self.dtype)
|
||||
self.current_depth = self.initial_depth
|
||||
depth_from = self.current_depth * self.depth_increment
|
||||
depth_to = depth_from + self.depth_increment
|
||||
if cm == "stacked":
|
||||
noise = self.noise_chunk[:, :, self.current_depth]
|
||||
noise = torch.cat(
|
||||
tuple(
|
||||
noise[bidx * self.batch : bidx * self.batch + self.batch]
|
||||
for bidx in range(noise.shape[0] // self.batch)
|
||||
),
|
||||
dim=2,
|
||||
)
|
||||
elif cm == "flat":
|
||||
noise = self.noise_chunk[:, :, self.current_depth]
|
||||
noise = noise.flatten()[: math.prod(self.shape)]
|
||||
elif cm in {"over_depth", "over_depth_alt"}:
|
||||
noise = self.noise_chunk[:, :, depth_from:depth_to]
|
||||
if cm == "over_depth":
|
||||
noise = noise.movedim(2, 1)
|
||||
elif cm == "over_depth_avg":
|
||||
noise = self.noise_chunk[:, :, depth_from:depth_to].mean(dim=1)
|
||||
elif cm.startswith("over_depth_"):
|
||||
channel_lookup = {"h": 0, "w": 1, "z": 2}
|
||||
mathop = cm[-5:-2]
|
||||
if mathop in {"add", "sub", "mul", "div"}:
|
||||
chan1, chan2 = channel_lookup[cm[-7]], channel_lookup[cm[-1]]
|
||||
noise1 = self.noise_chunk[:, chan1 : chan1 + 1, depth_from:depth_to]
|
||||
noise2 = self.noise_chunk[:, chan2 : chan2 + 1, depth_from:depth_to]
|
||||
if mathop == "sub":
|
||||
noise = noise1 - noise2
|
||||
elif mathop == "add":
|
||||
noise = noise1 + noise2
|
||||
elif mathop == "mul":
|
||||
noise = noise1 * noise2
|
||||
elif mathop == "div":
|
||||
noise = noise1 / (noise2 + 1e-07)
|
||||
else:
|
||||
chan = channel_lookup[cm[-1]]
|
||||
noise = self.noise_chunk[:, chan : chan + 1, depth_from:depth_to]
|
||||
else:
|
||||
raise ValueError("Bad channel mode")
|
||||
self.current_depth += 1 * self.direction
|
||||
if self.current_depth > self.max_depth or self.current_depth < 0:
|
||||
dm = self.depth_mode
|
||||
if dm == "reset":
|
||||
self.noise_chunk = None
|
||||
elif dm == "wrap":
|
||||
self.current_depth = self.initial_depth
|
||||
elif dm == "bounce":
|
||||
if self.depth < 2:
|
||||
raise ValueError("Bounce depth mode requires depth of at least 2")
|
||||
self.direction = -self.direction
|
||||
self.current_depth += 2 * self.direction
|
||||
return (
|
||||
noise.reshape(self.batch, -1, self.height, self.width)[
|
||||
:,
|
||||
: self.channels,
|
||||
]
|
||||
.clone()
|
||||
.contiguous()
|
||||
)
|
||||
@@ -0,0 +1,628 @@
|
||||
# ruff: noqa: ANN002, ANN003
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
|
||||
from .. import utils
|
||||
from .base import NoiseGenerator
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
|
||||
# With help from ChatGPT.
|
||||
class VoronoiNoiseGenerator(NoiseGenerator):
|
||||
name = "voronoi"
|
||||
MIN_DIMS = 4
|
||||
MAX_DIMS = 4
|
||||
|
||||
voronoi_distance_modes = frozenset((
|
||||
"angle_sigmoid",
|
||||
"angle_tanh",
|
||||
"angle",
|
||||
"chebyshev",
|
||||
"euclidean",
|
||||
"fractal_norm",
|
||||
"fuzz",
|
||||
"manhatten",
|
||||
"minkowski",
|
||||
"quadratic",
|
||||
"weight",
|
||||
))
|
||||
|
||||
voronoi_result_modes = frozenset((
|
||||
"cellid",
|
||||
"diff",
|
||||
"diff2",
|
||||
"f",
|
||||
"f1",
|
||||
"f2",
|
||||
"f3",
|
||||
"f4",
|
||||
"fractal_norm",
|
||||
"fuzz",
|
||||
"inv_f",
|
||||
"inv_f1",
|
||||
"inv_f2",
|
||||
"inv_f3",
|
||||
"inv_f4",
|
||||
"gradient_magnitude",
|
||||
"median_distance",
|
||||
"ridge",
|
||||
"softmin",
|
||||
))
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls, *, no_super: bool = False):
|
||||
result = {
|
||||
"n_points": (32,),
|
||||
"distance_mode": ("euclidean",),
|
||||
"z_initial": 0.0,
|
||||
"z_increment": 1.0,
|
||||
"z_max": 100000,
|
||||
"z_max_mode": "reset",
|
||||
# None or numeric
|
||||
"z_range": None,
|
||||
"result_mode": ("f1",),
|
||||
"octaves": 1,
|
||||
# same_features or new_features
|
||||
"octave_mode": "same_features",
|
||||
"lacunarity": 2.0, # scale increase per octave
|
||||
"gain": 0.5, # amplitude decrease per octave
|
||||
"initial_amplitude": 1.0,
|
||||
"initial_scale": 1.0,
|
||||
"noise_sampler_factory": None,
|
||||
"normalized": False,
|
||||
}
|
||||
return result if no_super else super().ng_params() | result
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.feature_points = self.grid_xyz = None
|
||||
self.noise_samplers = None
|
||||
self.n_points = tuple(max(2, val) for val in self.n_points)
|
||||
|
||||
def voronoi_reset(self, *args):
|
||||
self.z_curr = self.z_initial
|
||||
octave_range = tuple(
|
||||
range(self.octaves if self.octave_mode == "new_features" else 1),
|
||||
)
|
||||
if self.noise_sampler_factory is not None and self.noise_samplers is None:
|
||||
self.noise_samplers = tuple(
|
||||
self.noise_sampler_factory.make_noise_sampler(
|
||||
torch.zeros(
|
||||
self.batch,
|
||||
self.channels,
|
||||
self.n_points[octave % len(self.n_points)],
|
||||
3,
|
||||
device=self.gen_device,
|
||||
dtype=self.dtype,
|
||||
),
|
||||
cpu=self.cpu,
|
||||
normalized=False,
|
||||
)
|
||||
for octave in octave_range
|
||||
)
|
||||
self.feature_points = tuple(
|
||||
(
|
||||
torch.rand(
|
||||
self.batch,
|
||||
self.channels,
|
||||
self.n_points[octave % len(self.n_points)],
|
||||
3,
|
||||
device=self.gen_device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
if self.noise_samplers is None
|
||||
else utils.normalize_to_scale(
|
||||
self.noise_samplers[octave](*args),
|
||||
target_min=0.0,
|
||||
target_max=1.0,
|
||||
dim=(-1, -2),
|
||||
)
|
||||
).to(device=self.device)
|
||||
for octave in octave_range
|
||||
)
|
||||
if self.grid_xyz is not None:
|
||||
return
|
||||
y = torch.linspace(
|
||||
0,
|
||||
self.height - 1,
|
||||
self.height,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
x = torch.linspace(
|
||||
0,
|
||||
self.width - 1,
|
||||
self.width,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
self.grid_xyz = torch.stack(
|
||||
torch.meshgrid(y, x, indexing="ij"),
|
||||
dim=-1,
|
||||
) / torch.tensor(
|
||||
(self.height, self.width),
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
def get_feature_points(self, octave: int) -> torch.Tensor:
|
||||
result = self.feature_points[octave % len(self.feature_points)]
|
||||
odd_octave = (octave % 2) == 1
|
||||
om = self.octave_mode
|
||||
if (om == "same_invert_odd" and odd_octave) or (
|
||||
om == "same_invert_even" and not odd_octave
|
||||
):
|
||||
return 1.0 - result
|
||||
if octave > 0 and om in {"same_roll_chan_up", "same_roll_chan_down"}:
|
||||
return torch.roll(
|
||||
result,
|
||||
(-1 if om == "same_roll_chan_up" else 1) * (octave % 3),
|
||||
dims=(1,),
|
||||
)
|
||||
if octave > 0 and om in {"same_roll_dir_up", "same_roll_dir_down"}:
|
||||
return torch.roll(
|
||||
result,
|
||||
(-1 if om == "same_roll_dir_up" else 1) * (octave % 3),
|
||||
dims=(3,),
|
||||
)
|
||||
return result
|
||||
|
||||
def get_distance_mode(self, octave: int) -> torch.Tensor:
|
||||
return self.distance_mode[octave % len(self.distance_mode)]
|
||||
|
||||
def get_result_mode(self, octave: int) -> torch.Tensor:
|
||||
return self.result_mode[octave % len(self.result_mode)]
|
||||
|
||||
def voronoi_call_mode(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
result: bool,
|
||||
args: list | tuple = (),
|
||||
kwargs: dict | None = None,
|
||||
) -> torch.Tensor:
|
||||
name = name.strip().lower()
|
||||
modes = self.voronoi_result_modes if result else self.voronoi_distance_modes
|
||||
mode_label = "result" if result else "distance"
|
||||
if name not in modes:
|
||||
errstr = f"Bad Voronoi {mode_label} mode {name}"
|
||||
raise ValueError(errstr)
|
||||
kwargs = (
|
||||
{}
|
||||
if kwargs is None
|
||||
else {
|
||||
k[1:] if k.startswith("_") and len(k) > 1 else k: v
|
||||
for k, v in kwargs.items()
|
||||
}
|
||||
)
|
||||
return getattr(self, f"_voronoi_{mode_label}_{name}")(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_distance_euclidean(d: torch.Tensor, **_kwargs) -> torch.Tensor:
|
||||
return d.pow(2).sum(dim=-1).sqrt_()
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_distance_manhatten(d: torch.Tensor, **_kwargs) -> torch.Tensor:
|
||||
return d.pow(2).sum(dim=-1).sqrt_()
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_distance_chebyshev(d: torch.Tensor, **_kwargs) -> torch.Tensor:
|
||||
return d.abs().amax(dim=-1)
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_distance_minkowski(
|
||||
d: torch.Tensor,
|
||||
*,
|
||||
p: float | str = 3.0,
|
||||
**_kwargs,
|
||||
) -> torch.Tensor:
|
||||
p = float(p)
|
||||
return d.abs().pow(p).sum(dim=-1).pow(1 / p)
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_distance_quadratic(d: torch.Tensor, **_kwargs) -> torch.Tensor:
|
||||
return d.pow(2).sum(dim=-1)
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_distance_angle(
|
||||
d: torch.Tensor,
|
||||
*,
|
||||
idx: int | str = 2,
|
||||
**_kwargs,
|
||||
) -> torch.Tensor:
|
||||
return (
|
||||
torch.nn.functional.normalize(d, dim=-1)[..., int(idx)]
|
||||
.clamp_(-1.0, 1.0)
|
||||
.acos_()
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_distance_angle_tanh(
|
||||
d: torch.Tensor,
|
||||
*,
|
||||
idx: int | str = 2,
|
||||
**_kwargs,
|
||||
) -> torch.Tensor:
|
||||
return torch.nn.functional.normalize(d, dim=-1)[..., int(idx)].tanh_().acos_()
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_distance_angle_sigmoid(
|
||||
d: torch.Tensor,
|
||||
*,
|
||||
idx: int | str = 2,
|
||||
**_kwargs,
|
||||
) -> torch.Tensor:
|
||||
return (
|
||||
torch.nn.functional.normalize(d, dim=-1)[..., int(idx)]
|
||||
.sigmoid_()
|
||||
.mul_(2)
|
||||
.sub_(1)
|
||||
.acos_()
|
||||
)
|
||||
|
||||
def _voronoi_distance_weight(
|
||||
self,
|
||||
d: torch.Tensor,
|
||||
*args,
|
||||
name: str = "euclidean",
|
||||
h: float | str = 1.0,
|
||||
w: float | str = 1.0,
|
||||
z: float | str = 0.25,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
weights = d.new_tensor((float(h), float(w), float(z)))
|
||||
return self.voronoi_call_mode(
|
||||
name,
|
||||
result=False,
|
||||
args=(d * weights, *args),
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
def _voronoi_distance_fractal_norm(
|
||||
self,
|
||||
d: torch.Tensor,
|
||||
*args,
|
||||
name: str = "euclidean",
|
||||
mode: str = "sin",
|
||||
scale: str | float = 0.1,
|
||||
multiplier: str | float = 10.0,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if mode == "sin":
|
||||
fun = torch.sin
|
||||
elif mode == "cos":
|
||||
fun = torch.cos
|
||||
else:
|
||||
raise ValueError(
|
||||
"Bad mode parameter for fractal_norm distance mode, must be one of: sin, cos",
|
||||
)
|
||||
adjustment = float(scale) * fun(d * float(multiplier))
|
||||
return self.voronoi_call_mode(
|
||||
name,
|
||||
result=False,
|
||||
args=(d + adjustment, *args),
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
def _voronoi_distance_fuzz(
|
||||
self,
|
||||
*args,
|
||||
name: str = "euclidean",
|
||||
fuzz: float | str = 0.25,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
fuzz = float(fuzz)
|
||||
result = self.voronoi_call_mode(name, result=False, args=args, kwargs=kwargs)
|
||||
rmin, rmax = result.aminmax()
|
||||
fuzz = max(abs(rmin.item()), abs(rmax.item())) * fuzz
|
||||
result += (
|
||||
torch.rand(result.shape, device=self.gen_device, dtype=result.dtype)
|
||||
.mul_(fuzz * 2)
|
||||
.sub_(fuzz)
|
||||
.to(device=result.device)
|
||||
)
|
||||
return utils.normalize_to_scale(result, rmin.item(), rmax.item(), dim=(-2, -1))
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_result_f(
|
||||
_d: torch.Tensor,
|
||||
*,
|
||||
get_sorted: Callable,
|
||||
idx: int | str = 0,
|
||||
**_kwargs,
|
||||
) -> torch.Tensor:
|
||||
return get_sorted()[..., int(idx)]
|
||||
|
||||
def _voronoi_result_f1(self, *args, **kwargs) -> torch.Tensor:
|
||||
return self._voronoi_result_f(*args, idx=0, **kwargs)
|
||||
|
||||
def _voronoi_result_f2(self, *args, **kwargs) -> torch.Tensor:
|
||||
return self._voronoi_result_f(*args, idx=1, **kwargs)
|
||||
|
||||
def _voronoi_result_f3(self, *args, **kwargs) -> torch.Tensor:
|
||||
return self._voronoi_result_f(*args, idx=2, **kwargs)
|
||||
|
||||
def _voronoi_result_f4(self, *args, **kwargs) -> torch.Tensor:
|
||||
return self._voronoi_result_f(*args, idx=3, **kwargs)
|
||||
|
||||
def _voronoi_result_inv_f(self, *args, eps=1e-06, **kwargs) -> torch.Tensor:
|
||||
return 1.0 / (self._voronoi_result_f(*args, **kwargs) + eps)
|
||||
|
||||
def _voronoi_result_inv_f1(self, *args, **kwargs) -> torch.Tensor:
|
||||
return self._voronoi_result_inv_f(*args, idx=0, **kwargs)
|
||||
|
||||
def _voronoi_result_inv_f2(self, *args, **kwargs) -> torch.Tensor:
|
||||
return self._voronoi_result_inv_f(*args, idx=1, **kwargs)
|
||||
|
||||
def _voronoi_result_inv_f3(self, *args, **kwargs) -> torch.Tensor:
|
||||
return self._voronoi_result_inv_f(*args, idx=2, **kwargs)
|
||||
|
||||
def _voronoi_result_inv_f4(self, *args, **kwargs) -> torch.Tensor:
|
||||
return self._voronoi_result_inv_f(*args, idx=3, **kwargs)
|
||||
|
||||
def _voronoi_result_diff(
|
||||
self,
|
||||
*args,
|
||||
idx1: int | str = 0,
|
||||
idx2: int | str = 1,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
val1, val2 = (
|
||||
self._voronoi_result_f(*args, idx=i, **kwargs) for i in (idx1, idx2)
|
||||
)
|
||||
return val2 - val1
|
||||
|
||||
def _voronoi_result_diff2(
|
||||
self,
|
||||
*args,
|
||||
idx1: int | str = 0,
|
||||
idx2: int | str = 1,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
val1, val2 = (
|
||||
self._voronoi_result_f(*args, idx=i, **kwargs) for i in (idx1, idx2)
|
||||
)
|
||||
return (val2 - val1) / (val2 + val1 + 1e-06)
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_result_cellid(d, *_args, **_kwargs) -> torch.Tensor:
|
||||
cellids = d.argmin(dim=-1).to(dtype=d.dtype)
|
||||
return (cellids / cellids.max()).add_(1.0)
|
||||
|
||||
def _voronoi_result_ridge(
|
||||
self,
|
||||
*args,
|
||||
name: str = "diff",
|
||||
exp: float | str = -10.0,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
return 1.0 - (
|
||||
float(exp)
|
||||
* self.voronoi_call_mode(name, result=True, args=args, kwargs=kwargs)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_result_median_distance(
|
||||
*_args,
|
||||
get_sorted: Callable,
|
||||
**_kwargs,
|
||||
) -> torch.Tensor:
|
||||
return get_sorted().median(dim=-1).values
|
||||
|
||||
@staticmethod
|
||||
def _voronoi_result_softmin(
|
||||
d: torch.Tensor,
|
||||
*_args,
|
||||
temperature=50.0,
|
||||
use_sorted=None,
|
||||
d_orig: torch.Tensor,
|
||||
get_sorted: Callable,
|
||||
**_kwargs,
|
||||
) -> torch.Tensor:
|
||||
d_norm = d_orig.norm(dim=-1)
|
||||
soft_weights = F.softmax(-d_norm * float(temperature), dim=-1)
|
||||
eff_d = get_sorted() if use_sorted is not None else d
|
||||
return (eff_d * soft_weights).sum(dim=-1)
|
||||
|
||||
def _voronoi_result_gradient_magnitude(
|
||||
self,
|
||||
*args,
|
||||
name1: str = "f4",
|
||||
name2: str = "f4",
|
||||
pad_mode: str = "replicate",
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
r1 = self.voronoi_call_mode(name1, result=True, args=args, kwargs=kwargs)
|
||||
r1_padded = F.pad(r1, (1, 1, 1, 1), mode=pad_mode)
|
||||
if name2 != name1:
|
||||
r2 = self.voronoi_call_mode(name2, result=True, args=args, kwargs=kwargs)
|
||||
r2_padded = F.pad(r2, (1, 1, 1, 1), mode=pad_mode)
|
||||
else:
|
||||
r2 = r1
|
||||
r2_padded = r1_padded
|
||||
dx = r1_padded[..., 1:-1, 2:] - r2_padded[..., 1:-1, :-2]
|
||||
dy = r1_padded[..., 2:, 1:-1] - r2_padded[..., :-2, 1:-1]
|
||||
return (dx**2 + dy**2).sqrt_()
|
||||
|
||||
def _voronoi_result_fractal_norm(
|
||||
self,
|
||||
d: torch.Tensor,
|
||||
*args,
|
||||
name: str = "diff",
|
||||
mode: str = "sin",
|
||||
scale: str | float = 0.1,
|
||||
multiplier: str | float = 10.0,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if mode == "sin":
|
||||
fun = torch.sin
|
||||
elif mode == "cos":
|
||||
fun = torch.cos
|
||||
else:
|
||||
raise ValueError(
|
||||
"Bad mode parameter for fractal_norm result mode, must be one of: sin, cos",
|
||||
)
|
||||
d_adjusted = float(scale) * fun(d * float(multiplier))
|
||||
my_d_sorted = None
|
||||
|
||||
def my_get_sorted():
|
||||
nonlocal my_d_sorted
|
||||
if my_d_sorted is not None:
|
||||
return my_d_sorted
|
||||
my_d_sorted = d_adjusted.sort(dim=-1).values
|
||||
return my_d_sorted
|
||||
|
||||
return self.voronoi_call_mode(
|
||||
name,
|
||||
result=True,
|
||||
args=(d_adjusted, *args),
|
||||
kwargs=kwargs | {"get_sorted": my_get_sorted},
|
||||
)
|
||||
|
||||
def _voronoi_result_fuzz(
|
||||
self,
|
||||
*args,
|
||||
name: str = "f1",
|
||||
fuzz: float | str = 0.25,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
fuzz = float(fuzz)
|
||||
result = self.voronoi_call_mode(name, result=True, args=args, kwargs=kwargs)
|
||||
rmin, rmax = result.aminmax()
|
||||
fuzz = max(abs(rmin.item()), abs(rmax.item())) * fuzz
|
||||
result += (
|
||||
torch.rand(result.shape, device=self.gen_device, dtype=result.dtype)
|
||||
.mul_(fuzz * 2)
|
||||
.sub_(fuzz)
|
||||
.to(device=result.device)
|
||||
)
|
||||
return utils.normalize_to_scale(result, rmin.item(), rmax.item(), dim=(-2, -1))
|
||||
|
||||
def voronoi_distance(self, d: torch.Tensor, octave: int) -> torch.Tensor:
|
||||
modes = self.get_distance_mode(octave).split("+")
|
||||
result_scale_base = 1.0 / len(modes)
|
||||
result = None
|
||||
for mode in modes:
|
||||
if ":" in mode:
|
||||
mode_name, *mode_rest = mode.split(":")
|
||||
mode_kwargs = dict(
|
||||
tuple(val.strip() for val in di.split("=", 1)) for di in mode_rest
|
||||
)
|
||||
result_scale = result_scale_base * float(mode_kwargs.pop("dscale", 1.0))
|
||||
else:
|
||||
mode_name = mode
|
||||
mode_kwargs = {}
|
||||
result_scale = result_scale_base
|
||||
curr_result = self.voronoi_call_mode(
|
||||
mode_name,
|
||||
result=False,
|
||||
args=(d,),
|
||||
kwargs=mode_kwargs,
|
||||
).mul_(result_scale)
|
||||
result = curr_result if result is None else result.add_(curr_result)
|
||||
return result
|
||||
|
||||
def voronoi_result(
|
||||
self,
|
||||
d: torch.Tensor,
|
||||
d_orig: torch.Tensor,
|
||||
*,
|
||||
octave: int,
|
||||
) -> torch.Tensor:
|
||||
modes = self.get_result_mode(octave).split("+")
|
||||
result_scale_base = 1.0 / len(modes)
|
||||
result = None
|
||||
d_sorted = None
|
||||
|
||||
def get_sorted():
|
||||
nonlocal d_sorted
|
||||
if d_sorted is not None:
|
||||
return d_sorted
|
||||
d_sorted = d.sort(dim=-1).values
|
||||
return d_sorted
|
||||
|
||||
base_kwargs = {
|
||||
"d_orig": d_orig,
|
||||
"get_sorted": get_sorted,
|
||||
}
|
||||
for mode in modes:
|
||||
if ":" in mode:
|
||||
mode_name, *mode_rest = mode.split(":")
|
||||
mode_kwargs = dict(
|
||||
tuple(v.strip() for v in di.split("=", 1)) for di in mode_rest
|
||||
)
|
||||
result_scale = result_scale_base * float(mode_kwargs.pop("rscale", 1.0))
|
||||
else:
|
||||
result_scale = result_scale_base
|
||||
mode_name = mode
|
||||
mode_kwargs = {}
|
||||
curr_result = self.voronoi_call_mode(
|
||||
mode_name,
|
||||
result=True,
|
||||
args=(d,),
|
||||
kwargs=mode_kwargs | base_kwargs,
|
||||
).mul_(result_scale)
|
||||
result = curr_result if result is None else result.add_(curr_result)
|
||||
return result
|
||||
|
||||
def generate_octave(
|
||||
self,
|
||||
*,
|
||||
octave: int,
|
||||
grid: torch.Tensor,
|
||||
z_grid: torch.Tensor,
|
||||
scale: float = 1.0,
|
||||
) -> torch.Tensor:
|
||||
# Full 3D grid (H, W, 3)
|
||||
grid_3d = torch.cat((grid, z_grid), dim=-1)[None, None, ...] # (1, 1, H, W, 3)
|
||||
grid_3d = grid_3d.expand(self.batch, self.channels, -1, -1, -1)
|
||||
grid_3d = grid_3d.unsqueeze(-2) # (B, C, H, W, 1, 3)
|
||||
grid_3d = (grid_3d * scale) % 1.0
|
||||
|
||||
# Normalize feature points: already assumed in [0, 1)
|
||||
fp = self.get_feature_points(octave) # (B, C, N, 3)
|
||||
fp = fp[:, :, None, None] # (B, C, 1, 1, N, 3)
|
||||
fp = (fp * scale) % 1.0
|
||||
|
||||
# Toroidal wrapped difference
|
||||
d_orig = d = (grid_3d - fp + 0.5) % 1.0 - 0.5 # Wrap to [-0.5, 0.5)
|
||||
d = self.voronoi_distance(d.clone(), octave=octave)
|
||||
return self.voronoi_result(d, d_orig, octave=octave)
|
||||
|
||||
def generate(self, *args):
|
||||
if self.grid_xyz is None or self.feature_points is None or self.z_max == 0:
|
||||
self.voronoi_reset(*args)
|
||||
elif self.z_max != 0 and abs(self.z_initial - self.z_curr) > abs(self.z_max):
|
||||
if self.z_max_mode == "reset":
|
||||
self.voronoi_reset(*args)
|
||||
elif self.z_max_mode == "bounce":
|
||||
self.z_increment = -self.z_increment
|
||||
self.z_curr += self.z_increment
|
||||
else:
|
||||
self.curr_z = self.z_initial
|
||||
z_range = utils.fallback(self.z_range, max(self.height, self.width))
|
||||
z_norm = (self.z_curr % z_range) / z_range
|
||||
self.z_curr += self.z_increment
|
||||
grid = self.grid_xyz
|
||||
z_grid = grid.new_full((self.height, self.width, 1), z_norm)
|
||||
|
||||
result = grid.new_zeros(self.shape)
|
||||
amplitude = self.initial_amplitude
|
||||
scale = self.initial_scale
|
||||
total_amplitude = 0.0
|
||||
|
||||
for octave in range(self.octaves):
|
||||
result += self.generate_octave(
|
||||
octave=octave,
|
||||
grid=grid,
|
||||
z_grid=z_grid,
|
||||
scale=scale,
|
||||
).mul_(amplitude)
|
||||
total_amplitude += abs(amplitude)
|
||||
amplitude *= self.gain
|
||||
scale *= self.lacunarity
|
||||
result /= total_amplitude if total_amplitude != 0 else 1.0
|
||||
return result
|
||||
@@ -0,0 +1,138 @@
|
||||
# ruff: noqa: ANN002, ANN003
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from ..utils import fallback
|
||||
from ..wavelet_functions import Wavelet, wavelet_blend, wavelet_scaling
|
||||
from .base import FramesToChannelsNoiseGenerator
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
|
||||
# Idea from https://github.com/ClownsharkBatwing/RES4LYF/ (wave and mode defaults also from that source)
|
||||
class WaveletFilteredNoiseGenerator(FramesToChannelsNoiseGenerator):
|
||||
name = "waveletfilter"
|
||||
MIN_DIMS = 4
|
||||
MAX_DIMS = 5
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
inv_kwargs = {
|
||||
k: self.options[k]
|
||||
for k in ("inv_mode", "inv_biort", "inv_qshift", "inv_wave")
|
||||
if k in self.options
|
||||
}
|
||||
self.wavelet = Wavelet(
|
||||
wave=self.wave,
|
||||
level=self.level,
|
||||
mode=self.mode,
|
||||
use_1d_dwt=self.use_1d_dwt,
|
||||
use_dtcwt=self.use_dtcwt,
|
||||
biort=self.biort,
|
||||
qshift=self.qshift,
|
||||
device=self.gen_device,
|
||||
**inv_kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"mode": "periodization",
|
||||
"level": 3,
|
||||
"wave": "haar",
|
||||
"use_1d_dwt": False,
|
||||
"use_dtcwt": False,
|
||||
"qshift": "qshift_a",
|
||||
"biort": "near_sym_a",
|
||||
"yl_scale": 1.0,
|
||||
"yh_scales": 1.0,
|
||||
"two_step_inverse": False,
|
||||
"preblend_yl_scale_low": None,
|
||||
"preblend_yh_scales_low": None,
|
||||
"preblend_yl_scale_high": None,
|
||||
"preblend_yh_scales_high": None,
|
||||
"yl_blend_function": torch.lerp,
|
||||
"yh_blend_function": torch.lerp,
|
||||
"yl_blend_high": 0.0,
|
||||
"yh_blend_high": 1.0,
|
||||
"noise_sampler": None,
|
||||
"noise_sampler_high": None,
|
||||
}
|
||||
|
||||
def _fix_shape(self, noise, adjusted_shape):
|
||||
if noise.shape != adjusted_shape:
|
||||
noise = noise.reshape(*adjusted_shape)
|
||||
if self.frames:
|
||||
noise = noise.reshape(
|
||||
self.batch,
|
||||
self.channels * self.frames,
|
||||
self.height,
|
||||
self.width,
|
||||
)
|
||||
return noise
|
||||
|
||||
def generate(self, *args):
|
||||
adjusted_shape = self.get_adjusted_shape()
|
||||
noise = (
|
||||
self.rand_like()
|
||||
if self.noise_sampler is None
|
||||
else self.noise_sampler(*args)
|
||||
)
|
||||
if self.noise_sampler_high is not None:
|
||||
noise_high = self._fix_shape(self.noise_sampler_high(*args), adjusted_shape)
|
||||
else:
|
||||
noise_high = None
|
||||
noise = self._fix_shape(noise, adjusted_shape)
|
||||
orig_noise_shape = noise.shape
|
||||
need_flat = not self.use_dtcwt and self.use_1d_dwt and noise.ndim > 3
|
||||
if need_flat:
|
||||
noise = noise.flatten(start_dim=2)
|
||||
if noise_high is not None:
|
||||
noise_high = noise_high.flatten(start_dim=2)
|
||||
yl, yh = self.wavelet.forward(noise)
|
||||
if noise_high is not None:
|
||||
yl_high, yh_high = self.wavelet.forward(noise_high)
|
||||
if (
|
||||
self.preblend_yl_scale_high is not None
|
||||
or self.preblend_yh_scales_high is not None
|
||||
):
|
||||
yl_high, yh_high = wavelet_scaling(
|
||||
yl_high,
|
||||
yh_high,
|
||||
fallback(self.preblend_yl_scale_high, 1.0),
|
||||
fallback(self.preblend_yh_scales_high, 1.0),
|
||||
)
|
||||
if (
|
||||
self.preblend_yl_scale_low is not None
|
||||
or self.preblend_yh_scales_low is not None
|
||||
):
|
||||
yl, yh = wavelet_scaling(
|
||||
yl,
|
||||
yh,
|
||||
fallback(self.preblend_yl_scale_low, 1.0),
|
||||
fallback(self.preblend_yh_scales_low, 1.0),
|
||||
)
|
||||
yl, yh = wavelet_blend(
|
||||
(yl, yh),
|
||||
(yl_high, yh_high),
|
||||
yl_factor=self.yl_blend_high,
|
||||
yh_factor=self.yh_blend_high,
|
||||
blend_function=self.yl_blend_function,
|
||||
yh_blend_function=self.yh_blend_function,
|
||||
)
|
||||
del noise_high, yl_high, yh_high
|
||||
yl, yh = wavelet_scaling(
|
||||
yl,
|
||||
yh,
|
||||
self.yl_scale,
|
||||
self.yh_scales,
|
||||
in_place=True,
|
||||
)
|
||||
result = self.wavelet.inverse(yl, yh, two_step_inverse=self.two_step_inverse)
|
||||
if need_flat:
|
||||
result = result.reshape(orig_noise_shape)
|
||||
result = self.fix_output_frames(result)
|
||||
if result.shape == noise.shape:
|
||||
return result
|
||||
return result[tuple(slice(0, dl) for dl in noise.shape)]
|
||||
@@ -0,0 +1,146 @@
|
||||
# Some noise generation functions shamelessly yoinked from https://github.com/Clybius/ComfyUI-Extra-Samplers
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, NamedTuple
|
||||
|
||||
import torch
|
||||
|
||||
from .. import utils
|
||||
from .base import FramesToChannelsNoiseGenerator
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Sequence
|
||||
|
||||
|
||||
class WaveletNoiseOctave(NamedTuple):
|
||||
octave: int
|
||||
height: int
|
||||
width: int
|
||||
amplitude: float
|
||||
total_amplitude: float
|
||||
|
||||
|
||||
class WaveletNoiseGenerator(FramesToChannelsNoiseGenerator):
|
||||
name = "wavelet"
|
||||
MIN_DIMS = 4
|
||||
MAX_DIMS = 5
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"octave_scale_mode": "adaptive_avg_pool2d",
|
||||
"octave_rescale_mode": "bilinear",
|
||||
"post_octave_rescale_mode": "bilinear",
|
||||
"initial_amplitude": 1.0,
|
||||
"persistence": 0.5,
|
||||
"octaves": 4,
|
||||
"octave_height_factor": 0.5,
|
||||
"octave_width_factor": 0.5,
|
||||
"height_factor": 2.0,
|
||||
"width_factor": 2.0,
|
||||
"min_height": 4,
|
||||
"min_width": 4,
|
||||
"update_blend": 1.0,
|
||||
"update_blend_function": torch.lerp,
|
||||
"noise_sampler": None,
|
||||
}
|
||||
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.set_octave_data()
|
||||
|
||||
def set_internal_noise_sampler(self, noise_sampler: object) -> None:
|
||||
self.noise_sampler = noise_sampler
|
||||
|
||||
def set_octave_data(self) -> tuple:
|
||||
adjusted_shape = self.get_adjusted_shape()
|
||||
height, width = adjusted_shape[-2:]
|
||||
amplitude = self.initial_amplitude
|
||||
total_amplitude = 0.0
|
||||
curr_height, curr_width = height, width
|
||||
octave_data = []
|
||||
is_reverse = self.octaves < 0
|
||||
octaves = (
|
||||
range(self.octaves)
|
||||
if not is_reverse
|
||||
else reversed(range(abs(self.octaves)))
|
||||
)
|
||||
for octave in octaves:
|
||||
curr_height /= self.height_factor**octave
|
||||
curr_width /= self.width_factor**octave
|
||||
if (
|
||||
amplitude == 0
|
||||
or curr_height < self.min_height
|
||||
or curr_width < self.min_width
|
||||
or curr_height * self.octave_height_factor < 1
|
||||
or curr_width * self.octave_width_factor < 1
|
||||
):
|
||||
if is_reverse and not octave_data:
|
||||
curr_height, curr_width = height, width
|
||||
continue
|
||||
break
|
||||
total_amplitude += abs(amplitude)
|
||||
octave_data.append(
|
||||
WaveletNoiseOctave(
|
||||
octave=octave,
|
||||
height=curr_height,
|
||||
width=curr_width,
|
||||
amplitude=amplitude,
|
||||
total_amplitude=total_amplitude,
|
||||
),
|
||||
)
|
||||
amplitude *= self.persistence
|
||||
if not octave_data or not total_amplitude:
|
||||
raise ValueError("Unworkable parameters for wavelet noise")
|
||||
self.octave_data = tuple(octave_data)
|
||||
|
||||
def _generate_octave(self, *args: Any, shape: Sequence) -> torch.Tensor:
|
||||
height, width = shape[-2:]
|
||||
noise = (
|
||||
self.noise_sampler(*args)[..., :height, :width].reshape(shape)
|
||||
if self.noise_sampler
|
||||
else self.rand_like(shape=(*shape[:-2], height, width))
|
||||
)
|
||||
scaled_height = int(max(1, height * self.octave_height_factor))
|
||||
scaled_width = int(max(1, width * self.octave_width_factor))
|
||||
scaled_noise = utils.scale_samples(
|
||||
utils.scale_samples(
|
||||
noise,
|
||||
scaled_width,
|
||||
scaled_height,
|
||||
mode=self.octave_scale_mode,
|
||||
),
|
||||
width=width,
|
||||
height=height,
|
||||
mode=self.octave_rescale_mode,
|
||||
)
|
||||
return self.update_blend_function(
|
||||
noise,
|
||||
noise - scaled_noise,
|
||||
self.update_blend,
|
||||
)
|
||||
|
||||
def generate(self, *args: Any) -> torch.Tensor:
|
||||
adjusted_shape = self.get_adjusted_shape()
|
||||
height, width = adjusted_shape[-2:]
|
||||
curr_shape = list(adjusted_shape)
|
||||
result = torch.zeros(
|
||||
adjusted_shape,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
layout=self.layout,
|
||||
)
|
||||
for od in self.octave_data:
|
||||
curr_shape[-2:] = (int(od.height), int(od.width))
|
||||
octave_output = self._generate_octave(*args, shape=curr_shape)
|
||||
if octave_output.shape != result.shape:
|
||||
octave_output = utils.scale_samples(
|
||||
octave_output,
|
||||
width,
|
||||
height,
|
||||
mode=self.post_octave_rescale_mode,
|
||||
)
|
||||
result += octave_output.mul_(od.amplitude)
|
||||
if self.octave_data[-1].total_amplitude != 0:
|
||||
result /= self.octave_data[-1].total_amplitude
|
||||
return self.fix_output_frames(result)
|
||||
+34
-26
@@ -247,7 +247,7 @@ class SonarBase:
|
||||
momentum = self.cfg.momentum if momentum is None else momentum
|
||||
mode = self.cfg.momentum_mode
|
||||
if (
|
||||
momentum == 1 # noqa: PLR0916
|
||||
momentum == 1
|
||||
or history is None
|
||||
or (mode == MomentumMode.DENOISED and not is_denoised)
|
||||
or (mode != MomentumMode.DENOISED and is_denoised)
|
||||
@@ -368,39 +368,51 @@ class SonarGuidanceMixin:
|
||||
)
|
||||
raise ValueError("Sonar: Guidance: Unknown guidance type")
|
||||
|
||||
@staticmethod
|
||||
@classmethod
|
||||
def guidance_shift(cls, t: Tensor, ref_latent: Tensor, *, dim=None):
|
||||
if dim is None:
|
||||
dim = tuple(range(-(t.ndim - 1), 0))
|
||||
avg_t = t.mean(dim=dim, keepdim=True)
|
||||
std_t = t.std(dim=dim, keepdim=True)
|
||||
return (ref_latent * std_t).add_(avg_t)
|
||||
|
||||
@classmethod
|
||||
def guidance_euler(
|
||||
cls,
|
||||
sigma: Tensor,
|
||||
sigma_next: Tensor,
|
||||
x: Tensor,
|
||||
denoised: Tensor,
|
||||
ref_latent: Tensor,
|
||||
factor: float = 0.2,
|
||||
*,
|
||||
do_shift: bool = True,
|
||||
) -> Tensor:
|
||||
avg_t = denoised.mean(dim=(-3, -2, -1), keepdim=True)
|
||||
std_t = denoised.std(dim=(-3, -2, -1), keepdim=True)
|
||||
ref_img_shift = ref_latent * std_t + avg_t
|
||||
|
||||
if torch.equal(sigma, sigma_next):
|
||||
return cls.guidance_linear(x, ref_latent, factor=factor, do_shift=do_shift)
|
||||
ref_img_shift = (
|
||||
cls.guidance_shift(denoised, ref_latent) if do_shift else ref_latent
|
||||
)
|
||||
d = to_d(x, sigma, ref_img_shift)
|
||||
dt = (sigma_next - sigma) * factor
|
||||
return (d * dt).add_(x)
|
||||
|
||||
@staticmethod
|
||||
@classmethod
|
||||
def guidance_linear(
|
||||
cls,
|
||||
x: Tensor,
|
||||
ref_latent: Tensor,
|
||||
factor: float = 0.2,
|
||||
*,
|
||||
blend=torch.lerp,
|
||||
do_shift: bool = True,
|
||||
) -> Tensor:
|
||||
avg_t = x.mean(dim=(-3, -2, -1), keepdim=True)
|
||||
std_t = x.std(dim=(-3, -2, -1), keepdim=True)
|
||||
ref_img_shift = (ref_latent * std_t).add_(avg_t)
|
||||
ref_img_shift = cls.guidance_shift(x, ref_latent) if do_shift else ref_latent
|
||||
return blend(x, ref_img_shift, factor)
|
||||
|
||||
|
||||
class SonarWithGuidance(SonarBase, SonarGuidanceMixin):
|
||||
def __init__(self, *args: list[Any], **kwargs: dict[str, Any]):
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
SonarGuidanceMixin.__init__(self, self.cfg.guidance)
|
||||
|
||||
@@ -412,8 +424,8 @@ class SonarSampler(SonarWithGuidance):
|
||||
sigmas: Tensor,
|
||||
s_in: Tensor,
|
||||
extra_args: dict[str, Any],
|
||||
*args: list[Any],
|
||||
**kwargs: dict[str, Any],
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.model = model
|
||||
@@ -425,7 +437,7 @@ class SonarSampler(SonarWithGuidance):
|
||||
self,
|
||||
x: Tensor,
|
||||
sigma: Tensor,
|
||||
*args: list[Any],
|
||||
*args: Any,
|
||||
s_in=None,
|
||||
extra_args=None,
|
||||
) -> Tensor:
|
||||
@@ -440,8 +452,8 @@ class SonarSampler(SonarWithGuidance):
|
||||
class SonarEuler(SonarSampler):
|
||||
def __init__(
|
||||
self,
|
||||
*args: list[Any],
|
||||
**kwargs: dict[str, Any],
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
@@ -468,7 +480,6 @@ class SonarEuler(SonarSampler):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@torch.no_grad()
|
||||
def sampler(
|
||||
cls,
|
||||
model,
|
||||
@@ -520,8 +531,8 @@ class SonarEulerAncestral(SonarSampler):
|
||||
self,
|
||||
eta: float = 1.0,
|
||||
s_noise: float = 1.0,
|
||||
*args: list[Any],
|
||||
**kwargs: dict[str, Any],
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.eta = eta
|
||||
@@ -549,9 +560,8 @@ class SonarEulerAncestral(SonarSampler):
|
||||
)
|
||||
if sigma_next > 0:
|
||||
result_sample = self.guidance_step(step_index, result_sample, denoised)
|
||||
result_sample = ( # noqa: PLR6104
|
||||
result_sample
|
||||
+ self.noise_sampler(sigma, sigma_next) * (self.s_noise * sigma_up)
|
||||
result_sample = result_sample + self.noise_sampler(sigma, sigma_next) * (
|
||||
self.s_noise * sigma_up
|
||||
)
|
||||
|
||||
return (
|
||||
@@ -562,7 +572,6 @@ class SonarEulerAncestral(SonarSampler):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@torch.no_grad()
|
||||
def sampler(
|
||||
cls,
|
||||
model,
|
||||
@@ -620,8 +629,8 @@ class SonarDPMPPSDE(SonarSampler):
|
||||
self,
|
||||
eta: float = 1.0,
|
||||
s_noise: float = 1.0,
|
||||
*args: list[Any],
|
||||
**kwargs: dict[str, Any],
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.eta = eta
|
||||
@@ -760,7 +769,6 @@ class SonarDPMPPSDE(SonarSampler):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@torch.no_grad()
|
||||
def sampler(
|
||||
cls,
|
||||
model,
|
||||
|
||||
+1541
-43
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,846 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from enum import Enum, auto
|
||||
from typing import TYPE_CHECKING, Callable, NamedTuple
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from . import utils
|
||||
from .wavelet_functions import (
|
||||
Wavelet,
|
||||
expand_yh_scales,
|
||||
wavelet_blend,
|
||||
wavelet_scaling,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Sequence
|
||||
|
||||
|
||||
def pretty_non_default(obj: NamedTuple, *, defaults: object | None = None) -> str:
|
||||
result = ", ".join(
|
||||
f"{fn}={fv.pretty_non_default()}"
|
||||
if hasattr(fv, "pretty_non_default")
|
||||
else f"{fn}={fv!r}"
|
||||
for fn, fv in ((_fn, getattr(obj, _fn)) for _fn in obj._fields)
|
||||
if defaults is None or fv != getattr(defaults, fn)
|
||||
)
|
||||
return f"{obj.__class__.__name__}({result})"
|
||||
|
||||
|
||||
class WCFGSchedule(Enum):
|
||||
LINEAR = auto()
|
||||
LOGARITHMIC = auto()
|
||||
LOG = LOGARITHMIC
|
||||
EXPONENTIAL = auto()
|
||||
EXP = EXPONENTIAL
|
||||
HALF_COSINE = auto()
|
||||
SINE = auto()
|
||||
SIN = SINE
|
||||
|
||||
def interp(self, val: float) -> float:
|
||||
val = utils.clamp_float(val)
|
||||
if self == WCFGSchedule.LINEAR:
|
||||
return val
|
||||
if self == WCFGSchedule.LOGARITHMIC:
|
||||
result = 0.0 if val == 0 else math.log(val) + 1.0
|
||||
elif self == WCFGSchedule.EXPONENTIAL:
|
||||
result = math.exp(val) - 1.0
|
||||
elif self == WCFGSchedule.HALF_COSINE:
|
||||
result = 1.0 - ((1.0 + math.cos(val * math.pi)) / 2)
|
||||
elif self == WCFGSchedule.SINE:
|
||||
result = math.sin(val * math.pi)
|
||||
else:
|
||||
raise ValueError("Bad interpolation schedule!?")
|
||||
return utils.clamp_float(result)
|
||||
|
||||
|
||||
class WCFGSchedMode(Enum):
|
||||
SAMPLING = auto()
|
||||
ENABLED_SAMPLING = auto()
|
||||
SIGMAS = auto()
|
||||
ENABLED_SIGMAS = auto()
|
||||
STEP = auto()
|
||||
ENABLED_STEPS = auto()
|
||||
|
||||
# Aliases
|
||||
MODEL_SAMPLING = SAMPLING
|
||||
ENABLED_MODEL_SAMPLING = ENABLED_SAMPLING
|
||||
SIGMA_RANGE = SIGMAS
|
||||
ENABLED_SIGMA_RANGE = ENABLED_SIGMAS
|
||||
|
||||
|
||||
class WCFGTarget(Enum):
|
||||
DENOISED = auto()
|
||||
NOISE = auto()
|
||||
NOISE_NORM = auto()
|
||||
|
||||
|
||||
class WCFGPercentages(NamedTuple):
|
||||
sigma: float
|
||||
sigma_min: float
|
||||
sigma_max: float
|
||||
sigma_first: float | None
|
||||
sigma_last: float | None
|
||||
steps: int | None
|
||||
step: float | None
|
||||
step_first: int | None
|
||||
step_last: int | None
|
||||
pct_sampling: float
|
||||
pct_enabled_sampling: float
|
||||
pct_sigmas: float | None
|
||||
pct_enabled_sigmas: float | None
|
||||
pct_steps: float | None
|
||||
pct_enabled_steps: float | None
|
||||
|
||||
def invert(self) -> WCFGPercentages:
|
||||
return self._replace(
|
||||
pct_sampling=1.0 - self.pct_sampling,
|
||||
pct_enabled_sampling=1.0 - self.pct_enabled_sampling,
|
||||
pct_sigmas=None if self.pct_sigmas is None else 1.0 - self.pct_sigmas,
|
||||
pct_enabled_sigmas=None
|
||||
if self.pct_enabled_sigmas is None
|
||||
else 1.0 - self.pct_enabled_sigmas,
|
||||
pct_steps=None if self.pct_steps is None else 1.0 - self.pct_steps,
|
||||
pct_enabled_steps=None
|
||||
if self.pct_enabled_steps is None
|
||||
else 1.0 - self.pct_enabled_steps,
|
||||
)
|
||||
|
||||
def pct_from_schedmode(self, mode: WCFGSchedMode) -> float | None:
|
||||
if mode == WCFGSchedMode.MODEL_SAMPLING:
|
||||
return self.pct_sampling
|
||||
if mode == WCFGSchedMode.SIGMA_RANGE:
|
||||
return self.pct_sigmas
|
||||
if mode == WCFGSchedMode.ENABLED_MODEL_SAMPLING:
|
||||
return self.pct_enabled_sampling
|
||||
if mode == WCFGSchedMode.ENABLED_SIGMA_RANGE:
|
||||
return self.pct_enabled_sigmas
|
||||
if mode == WCFGSchedMode.STEP:
|
||||
if self.pct_steps is None:
|
||||
raise RuntimeError("Step percentage not available")
|
||||
return self.pct_steps
|
||||
raise ValueError("Unknown mode")
|
||||
|
||||
@classmethod
|
||||
def build(
|
||||
cls,
|
||||
*,
|
||||
ms: object,
|
||||
start_sigma: float,
|
||||
end_sigma: float,
|
||||
sigma: float,
|
||||
sigmas: torch.Tensor | None,
|
||||
**_kwargs: dict,
|
||||
) -> WCFGPercentages:
|
||||
if start_sigma < end_sigma:
|
||||
raise ValueError("start/end sigmas out of order")
|
||||
sigma_max = ms.sigma_max.detach().item()
|
||||
sigma_min = ms.sigma_min.detach().item()
|
||||
start_sigma = min(sigma_max, start_sigma)
|
||||
end_sigma = min(max(sigma_min, end_sigma), sigma_max)
|
||||
sigma = min(max(sigma, sigma_min), sigma_max)
|
||||
rstart = torch.tensor(start_sigma)
|
||||
rend = torch.tensor(end_sigma)
|
||||
pct_start = 1.0 - (ms.timestep(rstart) / 999).clamp(0, 1).detach().item()
|
||||
pct_end = 1.0 - (ms.timestep(rend) / 999).clamp(0, 1).detach().item()
|
||||
pct_curr = (
|
||||
1.0 - (ms.timestep(torch.tensor(sigma)) / 999).clamp(0, 1).detach().item()
|
||||
)
|
||||
pct_range_curr = (pct_curr - pct_start) / (pct_end - pct_start)
|
||||
|
||||
if sigmas is not None:
|
||||
if sigmas.ndim == 2:
|
||||
sigmas = sigmas.max(dim=0).values
|
||||
elif sigmas.ndim != 1:
|
||||
raise ValueError("Unexpected number of dimensions for sample_sigmas")
|
||||
sigmas = sigmas.detach().cpu()
|
||||
sigma_first = sigmas[0].item()
|
||||
sigma_last = sigmas[-2].item()
|
||||
if sigma_first <= sigma_last:
|
||||
raise ValueError(
|
||||
"Cannot handle non-descending sigmas (possibly Restart or unsampling)",
|
||||
)
|
||||
pct_sigmas = (sigma_first - sigma) / (sigma_first - sigma_last)
|
||||
start_sigma = min(start_sigma, sigma_first)
|
||||
end_sigma = max(end_sigma, sigma_last)
|
||||
sigma = min(max(sigma, sigma_last), sigma_first)
|
||||
if start_sigma == end_sigma:
|
||||
pct_enabled_sigmas = 1.0
|
||||
else:
|
||||
pct_enabled_sigmas = (start_sigma - sigma) / (start_sigma - end_sigma)
|
||||
steps = len(sigmas) - 1
|
||||
have_steps = False
|
||||
if steps > 1:
|
||||
step = utils.step_from_sigmas(sigma, sigmas)
|
||||
pct_steps = step / (steps - 1) if step is not None else None
|
||||
enabled_steps = torch.arange(len(sigmas), dtype=torch.int32)[
|
||||
(sigmas <= start_sigma) & (sigmas >= end_sigma)
|
||||
]
|
||||
if len(enabled_steps) > 1:
|
||||
have_steps = True
|
||||
step_first = enabled_steps[0].item()
|
||||
step_last = enabled_steps[-1].item()
|
||||
pct_enabled_steps = (step - step_first) / (step_last - step_first)
|
||||
if not have_steps:
|
||||
step = 0.0
|
||||
pct_steps = 1.0
|
||||
step_first = step_last = None
|
||||
pct_enabled_steps = None
|
||||
else:
|
||||
pct_enabled_sigmas = pct_sigmas = None
|
||||
step = steps = None
|
||||
pct_enabled_steps = pct_steps = None
|
||||
sigma_first = sigma_last = None
|
||||
return WCFGPercentages(
|
||||
pct_sampling=pct_curr,
|
||||
pct_enabled_sampling=pct_range_curr,
|
||||
pct_sigmas=pct_sigmas,
|
||||
pct_enabled_sigmas=pct_enabled_sigmas,
|
||||
pct_steps=pct_steps,
|
||||
pct_enabled_steps=pct_enabled_steps,
|
||||
sigma=sigma,
|
||||
sigma_first=sigma_first,
|
||||
sigma_last=sigma_last,
|
||||
sigma_min=sigma_min,
|
||||
sigma_max=sigma_max,
|
||||
steps=steps,
|
||||
step=step,
|
||||
step_first=step_first,
|
||||
step_last=step_last,
|
||||
)
|
||||
|
||||
|
||||
class WCFGScales(NamedTuple):
|
||||
yl_scale: float = 1.0
|
||||
yh_scales: float | Sequence = 1.0
|
||||
|
||||
def get_scales(
|
||||
self,
|
||||
*_args: list,
|
||||
verbose: bool = False,
|
||||
**_kwargs: dict,
|
||||
) -> WCFGScales:
|
||||
if verbose:
|
||||
tqdm.write(f"WCFG: {self.pretty_scales()}")
|
||||
return self
|
||||
|
||||
def apply_scales(
|
||||
self,
|
||||
yl: torch.Tensor,
|
||||
yh: Sequence,
|
||||
) -> tuple[torch.Tensor, Sequence]:
|
||||
return wavelet_scaling(yl, yh, yl_scale=self.yl_scale, yh_scales=self.yh_scales)
|
||||
|
||||
def get_and_apply_scales(
|
||||
self,
|
||||
pcts: WCFGPercentages,
|
||||
yl: torch.Tensor,
|
||||
yh: Sequence,
|
||||
*,
|
||||
verbose: bool = False,
|
||||
) -> tuple[torch.Tensor, Sequence]:
|
||||
return self.get_scales(pcts, yh, verbose=verbose).apply_scales(yl, yh)
|
||||
|
||||
def pretty_yh_scales(self, *, target=None) -> str:
|
||||
if target is None:
|
||||
target = self.yh_scales
|
||||
if isinstance(target, float):
|
||||
return f"{target:.4f}"
|
||||
if not isinstance(target, (list, tuple)):
|
||||
return str(target)
|
||||
result = ", ".join(
|
||||
self.pretty_yh_scales(target=val)
|
||||
if isinstance(val, (list, tuple))
|
||||
else (val if isinstance(val, str) else f"{val:.4f}")
|
||||
for val in target
|
||||
)
|
||||
return f"({result})"
|
||||
|
||||
def pretty_scales(self):
|
||||
return f"low={self.yl_scale:.4f}, high={self.pretty_yh_scales()}"
|
||||
|
||||
|
||||
class WCFGScheduledScale(NamedTuple):
|
||||
schedule: WCFGSchedule = WCFGSchedule.LINEAR
|
||||
schedule_mode: WCFGSchedMode = WCFGSchedMode.ENABLED_MODEL_SAMPLING
|
||||
schedule_offset: float = 0.0
|
||||
schedule_offset_after: float = 0.0
|
||||
schedule_multiplier: float = 1.0
|
||||
schedule_multiplier_after: float = 1.0
|
||||
reverse_schedule: bool = False
|
||||
reverse_schedule_after: bool = False
|
||||
schedule_min: float = 0.0
|
||||
schedule_max: float = 1.0
|
||||
|
||||
@classmethod
|
||||
def build(cls, **kwargs: dict) -> WCFGScheduledScale:
|
||||
schedule = kwargs.pop("schedule", DEFAULT_SCHEDULEDSCALE.schedule)
|
||||
if isinstance(schedule, str):
|
||||
schedule = getattr(WCFGSchedule, schedule.upper())
|
||||
schedule_mode = kwargs.pop(
|
||||
"schedule_mode",
|
||||
DEFAULT_SCHEDULEDSCALE.schedule_mode,
|
||||
)
|
||||
if isinstance(schedule_mode, str):
|
||||
schedule_mode = getattr(WCFGSchedMode, schedule_mode.upper())
|
||||
return WCFGScheduledScale(
|
||||
schedule=schedule,
|
||||
schedule_mode=schedule_mode,
|
||||
**utils.filter_dict(kwargs, cls._fields),
|
||||
)
|
||||
|
||||
def get_b_scale(self, pcts: WCFGPercentages) -> float:
|
||||
if self.reverse_schedule:
|
||||
pcts = pcts.invert()
|
||||
pct = pcts.pct_from_schedmode(self.schedule_mode)
|
||||
if pct is None:
|
||||
raise RuntimeError("Couldn't get percentage")
|
||||
pct = utils.clamp_float(
|
||||
(
|
||||
self.schedule.interp(
|
||||
utils.clamp_float(
|
||||
(pct + self.schedule_offset) * self.schedule_multiplier,
|
||||
),
|
||||
)
|
||||
+ self.schedule_offset_after
|
||||
)
|
||||
* self.schedule_multiplier_after,
|
||||
minval=utils.clamp_float(self.schedule_min),
|
||||
maxval=utils.clamp_float(self.schedule_max),
|
||||
)
|
||||
if self.reverse_schedule_after:
|
||||
pct = utils.clamp_float(1.0 - pct)
|
||||
return pct
|
||||
|
||||
def pretty_non_default(self) -> str:
|
||||
return pretty_non_default(self, defaults=DEFAULT_SCHEDULEDSCALE)
|
||||
|
||||
|
||||
DEFAULT_SCHEDULEDSCALE = WCFGScheduledScale()
|
||||
|
||||
|
||||
class WCFGScalesRange(NamedTuple):
|
||||
scales_start: WCFGScales = WCFGScales()
|
||||
scales_end: WCFGScales | None = None
|
||||
scheduler: WCFGScheduledScale | None = None
|
||||
blend_mode: str = "lerp"
|
||||
|
||||
@classmethod
|
||||
def build(cls, **kwargs: dict) -> WCFGScales | WCFGScalesRange:
|
||||
scales_start = kwargs.pop("scales_start", None)
|
||||
if scales_start is None:
|
||||
scales_start = {
|
||||
"yl_scale": kwargs.pop("yl_scale", 1.0),
|
||||
"yh_scales": kwargs.pop("yh_scales", 1.0),
|
||||
}
|
||||
scales_end = utils.filter_dict(kwargs.pop("scales_end", {}), WCFGScales._fields)
|
||||
if not scales_end or scales_end == scales_start:
|
||||
return WCFGScales(
|
||||
yl_scale=scales_start.get("yl_scale", 1.0),
|
||||
yh_scales=scales_start.get("yh_scales", 1.0),
|
||||
)
|
||||
blend_mode = kwargs.pop("blend_mode", "lerp")
|
||||
return WCFGScalesRange(
|
||||
scales_start=WCFGScales(**scales_start),
|
||||
scales_end=WCFGScales(**scales_end),
|
||||
scheduler=utils.maybe_apply_kwargs(
|
||||
kwargs,
|
||||
bool(scales_end),
|
||||
WCFGScheduledScale.build,
|
||||
),
|
||||
blend_mode=blend_mode,
|
||||
)
|
||||
|
||||
def get_scales(
|
||||
self,
|
||||
pcts: WCFGPercentages,
|
||||
yh: Sequence,
|
||||
*,
|
||||
verbose: bool = False,
|
||||
) -> WCFGScales:
|
||||
if self.scales_end is None or self.scheduler is None:
|
||||
return self.scales_start.get_scales()
|
||||
pct = self.scheduler.get_b_scale(pcts)
|
||||
if verbose:
|
||||
tqdm.write(f"WCFG: pct={pct:.4f}, percentages: {pcts}")
|
||||
start, end = self.scales_start, self.scales_end
|
||||
simple_blend = self.blend_mode == "lerp"
|
||||
if pct <= 0 and simple_blend:
|
||||
simple_result = start
|
||||
elif pct >= 1 and simple_blend:
|
||||
simple_result = end
|
||||
else:
|
||||
simple_result = None
|
||||
if simple_result is not None:
|
||||
if verbose:
|
||||
tqdm.write(
|
||||
f"WCFG: {simple_result.pretty_scales()}",
|
||||
)
|
||||
return simple_result
|
||||
start_yh_scales = expand_yh_scales(yh, yh_scales=start.yh_scales)
|
||||
end_yh_scales = expand_yh_scales(yh, yh_scales=end.yh_scales)
|
||||
blend_function = (
|
||||
None if self.blend_mode == "lerp" else utils.BLENDING_MODES[self.blend_mode]
|
||||
)
|
||||
yl_scale = utils.blend_scalar(
|
||||
start.yl_scale,
|
||||
end.yl_scale,
|
||||
pct,
|
||||
blend_function=blend_function,
|
||||
)
|
||||
yh_scales = tuple(
|
||||
tuple(
|
||||
utils.blend_scalar(os, oe, pct, blend_function=blend_function)
|
||||
for os, oe in zip(bs, be)
|
||||
)
|
||||
for bs, be in zip(start_yh_scales, end_yh_scales)
|
||||
)
|
||||
result = WCFGScales(yl_scale=yl_scale, yh_scales=yh_scales)
|
||||
if verbose:
|
||||
tqdm.write(
|
||||
f"WCFG: {result.pretty_scales()}",
|
||||
)
|
||||
return result
|
||||
|
||||
def apply_scales(
|
||||
self,
|
||||
yl: torch.Tensor,
|
||||
yh: Sequence,
|
||||
) -> tuple[torch.Tensor, Sequence]:
|
||||
return self.scales_start.apply_scales(yl, yh)
|
||||
|
||||
def get_and_apply_scales(
|
||||
self,
|
||||
pcts: WCFGPercentages,
|
||||
yl: torch.Tensor,
|
||||
yh: Sequence,
|
||||
*,
|
||||
verbose: bool = False,
|
||||
) -> tuple[torch.Tensor, Sequence]:
|
||||
return self.get_scales(pcts, yh, verbose=verbose).apply_scales(yl, yh)
|
||||
|
||||
def pretty_non_default(self) -> str:
|
||||
return pretty_non_default(self, defaults=DEFAULT_SCALESRANGE)
|
||||
|
||||
|
||||
DEFAULT_SCALESRANGE = WCFGScalesRange()
|
||||
|
||||
|
||||
class WCFGScheduledFloat(NamedTuple):
|
||||
value_start: float
|
||||
value_end: float | None = None
|
||||
scheduler: WCFGScheduledScale | None = None
|
||||
|
||||
@classmethod
|
||||
def build(
|
||||
cls,
|
||||
val: float | dict,
|
||||
*,
|
||||
default_start: float | None = None,
|
||||
default_end: float | None = None,
|
||||
**_kwargs: dict,
|
||||
) -> WCFGScheduledFloat:
|
||||
if isinstance(val, float):
|
||||
return WCFGScheduledFloat(value_start=val)
|
||||
if not isinstance(val, dict):
|
||||
raise TypeError("Bad type for scheduled float value")
|
||||
val = val.copy()
|
||||
value_start = val.pop("value_start", default_start)
|
||||
value_end = val.pop("value_end", default_end)
|
||||
if not isinstance(value_start, (float, int)):
|
||||
raise TypeError("Bad type for scheduled float start_value")
|
||||
if value_end is None:
|
||||
return WCFGScheduledFloat(value_start=val)
|
||||
if not isinstance(value_end, (float, int)):
|
||||
raise TypeError("Bad type for scheduled float end_value")
|
||||
return WCFGScheduledFloat(
|
||||
value_start=float(value_start),
|
||||
value_end=float(value_end),
|
||||
scheduler=WCFGScheduledScale.build(**val),
|
||||
)
|
||||
|
||||
def get_value(self, pcts: WCFGPercentages) -> float:
|
||||
if self.value_end is None or self.scheduler is None:
|
||||
return self.value_start
|
||||
pct = self.scheduler.get_b_scale(pcts)
|
||||
return (1.0 - pct) * self.value_start + pct * self.value_end
|
||||
|
||||
|
||||
class WCFGWaveletSettings(NamedTuple):
|
||||
wave: str = "db4"
|
||||
level: int = 5
|
||||
padding_mode: str = "symmetric"
|
||||
use_1d_dwt: bool = False
|
||||
use_dtcwt: bool = False
|
||||
biort: str = "near_sym_a"
|
||||
qshift: str = "qshift_a"
|
||||
inv_wave: str | None = None
|
||||
inv_padding_mode: str | None = None
|
||||
inv_biort: str | None = None
|
||||
inv_qshift: str | None = None
|
||||
|
||||
@classmethod
|
||||
def build(cls, **kwargs: dict) -> WCFGWaveletSettings:
|
||||
return WCFGWaveletSettings(**utils.filter_dict(kwargs, cls._fields))
|
||||
|
||||
def make_wavelet(self, **kwargs: dict) -> Wavelet:
|
||||
return Wavelet(
|
||||
wave=self.wave,
|
||||
level=self.level,
|
||||
mode=self.padding_mode,
|
||||
use_1d_dwt=self.use_1d_dwt,
|
||||
use_dtcwt=self.use_dtcwt,
|
||||
biort=self.biort,
|
||||
qshift=self.qshift,
|
||||
inv_wave=self.inv_wave,
|
||||
inv_mode=self.inv_padding_mode,
|
||||
inv_biort=self.inv_biort,
|
||||
inv_qshift=self.inv_qshift,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def pretty_non_default(self) -> str:
|
||||
return pretty_non_default(self, defaults=DEFAULT_WAVELETSETTINGS)
|
||||
|
||||
|
||||
DEFAULT_WAVELETSETTINGS = WCFGWaveletSettings()
|
||||
|
||||
|
||||
class WCFGRule(NamedTuple):
|
||||
start_sigma: float = math.inf
|
||||
end_sigma: float = 0.0
|
||||
verbose: bool = False
|
||||
blend_mode: str = "lerp"
|
||||
blend_strength: WCFGScheduledFloat = WCFGScheduledFloat(1.0)
|
||||
fallback_existing: bool = True
|
||||
target_mode: WCFGTarget = WCFGTarget.DENOISED
|
||||
diff: WCFGScalesRange | WCFGScales | None = None
|
||||
cond: WCFGScalesRange | WCFGScales | None = None
|
||||
uncond: WCFGScalesRange | WCFGScales | None = None
|
||||
final: WCFGScalesRange | WCFGScales | None = None
|
||||
wavelet: WCFGWaveletSettings = DEFAULT_WAVELETSETTINGS
|
||||
high_precision_mode: bool = True
|
||||
difference_blend_mode: str = "inject"
|
||||
difference_blend_strength: WCFGScheduledFloat = WCFGScheduledFloat(1.0)
|
||||
|
||||
@classmethod
|
||||
def build(cls, **kwargs: dict) -> WCFGRule:
|
||||
target_mode = kwargs.pop("target_mode", DEFAULT_RULE.target_mode)
|
||||
if isinstance(target_mode, str):
|
||||
target_mode = getattr(WCFGTarget, target_mode.upper())
|
||||
difference = kwargs.pop("diff", None)
|
||||
if difference is None:
|
||||
difference = kwargs.pop("difference", None)
|
||||
if difference is not None:
|
||||
difference = WCFGScalesRange.build(**difference)
|
||||
cond = kwargs.pop("cond", None)
|
||||
if cond is not None:
|
||||
cond = WCFGScalesRange.build(**cond)
|
||||
uncond = kwargs.pop("uncond", None)
|
||||
if uncond is not None:
|
||||
uncond = WCFGScalesRange.build(**uncond)
|
||||
final = kwargs.pop("final", None)
|
||||
if final is not None:
|
||||
final = WCFGScalesRange.build(**final)
|
||||
blend_strength = kwargs.pop("blend_strength", 1.0)
|
||||
if not isinstance(blend_strength, (float, int, dict)):
|
||||
raise TypeError("Bad type for blend_strength, must be float or dict")
|
||||
difference_blend_strength = kwargs.pop("difference_blend_strength", 1.0)
|
||||
if not isinstance(difference_blend_strength, (float, int, dict)):
|
||||
raise TypeError(
|
||||
"Bad type for difference_blend_strength, must be float or dict",
|
||||
)
|
||||
return WCFGRule(
|
||||
target_mode=target_mode,
|
||||
diff=difference,
|
||||
cond=cond,
|
||||
uncond=uncond,
|
||||
final=final,
|
||||
blend_strength=WCFGScheduledFloat(blend_strength),
|
||||
difference_blend_strength=WCFGScheduledFloat(difference_blend_strength),
|
||||
wavelet=WCFGWaveletSettings.build(**kwargs),
|
||||
**utils.filter_dict(kwargs, cls._fields),
|
||||
)
|
||||
|
||||
def make_wavelet(self, **kwargs: dict) -> Wavelet:
|
||||
return self.wavelet.make_wavelet(**kwargs)
|
||||
|
||||
def get_and_apply_scales(
|
||||
self,
|
||||
name: str,
|
||||
pcts: WCFGPercentages,
|
||||
yl: torch.Tensor,
|
||||
yh: Sequence,
|
||||
*,
|
||||
verbose: bool = False,
|
||||
) -> tuple[torch.Tensor, Sequence]:
|
||||
scales = getattr(self, name).get_scales(pcts, yh)
|
||||
if verbose and (scales.yl_scale != 1.0 or scales.yh_scales != 1.0):
|
||||
tqdm.write(
|
||||
f"WCFG: scales({name:>6}): {scales.pretty_scales()}",
|
||||
)
|
||||
return scales.apply_scales(yl, yh)
|
||||
|
||||
def pretty_non_default(self) -> str:
|
||||
return pretty_non_default(self, defaults=DEFAULT_RULE)
|
||||
|
||||
|
||||
DEFAULT_RULE = WCFGRule()
|
||||
|
||||
|
||||
class WCFGRules(NamedTuple):
|
||||
rules: Sequence = ()
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.rules)
|
||||
|
||||
def __getitem__(self, idx: int) -> WCFGRule:
|
||||
return self.rules[idx]
|
||||
|
||||
def __bool__(self) -> bool:
|
||||
return bool(self.rules)
|
||||
|
||||
def get_rule(self, sigma: float) -> WCFGRule | None:
|
||||
for rule in self.rules:
|
||||
if (
|
||||
rule.end_sigma
|
||||
<= sigma
|
||||
<= (math.inf if rule.start_sigma < 0 else rule.start_sigma)
|
||||
):
|
||||
return rule
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def build(cls, **params: dict) -> WCFGRules:
|
||||
params = params.copy()
|
||||
rules = params.pop("rules", ())
|
||||
rule_1 = WCFGRule.build(**params)
|
||||
other_rules = (WCFGRule.build(**rparams) for rparams in rules)
|
||||
return WCFGRules(rules=(rule_1, *other_rules))
|
||||
|
||||
|
||||
class WCFGContext(NamedTuple):
|
||||
cond: torch.Tensor
|
||||
uncond: torch.Tensor
|
||||
x: torch.Tensor
|
||||
sigma: torch.Tensor
|
||||
wavelet: Wavelet
|
||||
dtype: torch.dtype
|
||||
op_kwargs: dict
|
||||
|
||||
|
||||
class WaveletCFG:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
existing_cfg: Callable | None,
|
||||
rules: WCFGRules,
|
||||
operation_cond: Callable | None = None,
|
||||
operation_uncond: Callable | None = None,
|
||||
operation_fallback_cfg: Callable | None = None,
|
||||
operation_wavelet_cfg: Callable | None = None,
|
||||
operation_result: Callable | None = None,
|
||||
):
|
||||
self.wavelet_cache = {}
|
||||
self.rules = rules
|
||||
self.fallback_cfg_function = (
|
||||
existing_cfg
|
||||
if existing_cfg is not None and (not rules or rules[0].fallback_existing)
|
||||
else self.basic_cfg_function
|
||||
)
|
||||
self.operation_cond = operation_cond
|
||||
self.operation_uncond = operation_uncond
|
||||
self.operation_fallback_cfg = operation_fallback_cfg
|
||||
self.operation_wavelet_cfg = operation_wavelet_cfg
|
||||
self.operation_result = operation_result
|
||||
|
||||
@staticmethod
|
||||
def basic_cfg_function(args: dict) -> torch.Tensor:
|
||||
x, scale = args["input"], args["cond_scale"]
|
||||
uncond, cond = args["uncond_denoised"], args["cond_denoised"]
|
||||
return x - (cond - uncond).mul_(scale).add_(uncond)
|
||||
|
||||
@staticmethod
|
||||
def maybe_op(
|
||||
t: torch.Tensor,
|
||||
mop: Callable | None,
|
||||
**kwargs: dict,
|
||||
) -> torch.Tensor:
|
||||
return (
|
||||
t
|
||||
if mop is None
|
||||
else mop(
|
||||
latent=t,
|
||||
**(kwargs if getattr(mop, "EXTENDED_LATENT_OPERATION", None) else {}),
|
||||
)
|
||||
)
|
||||
|
||||
def get_context(self, *, rule: WCFGRule, args: dict) -> WCFGContext:
|
||||
sigma_orig = sigma = args["sigma"]
|
||||
rule_id = id(rule)
|
||||
x = args["input"]
|
||||
if x.ndim == 3 and not rule.wavelet.use_1d_dwt:
|
||||
raise RuntimeError("Enable use_1d_dwt mode for 3D latents.")
|
||||
if x.ndim < 3:
|
||||
raise RuntimeError(
|
||||
"Wavelet CFG can't handle latents with 2 or less dimensions.",
|
||||
)
|
||||
if sigma.ndim != x.ndim:
|
||||
sigma = sigma.reshape(x.shape[0], *((1,) * (x.ndim - sigma.ndim)))
|
||||
if rule.target_mode in {WCFGTarget.NOISE, WCFGTarget.NOISE_NORM}:
|
||||
cond, uncond = args["cond"], args["uncond"]
|
||||
if rule.target_mode == WCFGTarget.NOISE_NORM:
|
||||
cond = cond / sigma # noqa: PLR6104
|
||||
uncond = uncond / sigma # noqa: PLR6104
|
||||
elif rule.target_mode == WCFGTarget.DENOISED:
|
||||
cond, uncond = args["cond_denoised"], args["uncond_denoised"]
|
||||
else:
|
||||
raise ValueError("Bad target mode")
|
||||
op_kwargs = {
|
||||
"sigma": sigma_orig,
|
||||
"cond": cond,
|
||||
"uncond": uncond,
|
||||
"cond_scale": args["cond_scale"],
|
||||
"raw_args": args,
|
||||
}
|
||||
cond = self.maybe_op(cond, self.operation_cond, **op_kwargs)
|
||||
uncond = self.maybe_op(uncond, self.operation_uncond, **op_kwargs)
|
||||
eff_dtype = torch.float64 if rule.high_precision_mode else x.dtype
|
||||
wavelet = self.wavelet_cache.get(rule_id)
|
||||
if wavelet is None:
|
||||
wavelet = rule.make_wavelet()
|
||||
self.wavelet_cache[rule_id] = wavelet
|
||||
wavelet = wavelet.to(device=x.device, dtype=eff_dtype)
|
||||
if rule.wavelet.use_1d_dwt:
|
||||
cond = cond.flatten(start_dim=2)
|
||||
uncond = uncond.flatten(start_dim=2)
|
||||
elif x.ndim > 4:
|
||||
cond = cond.flatten(start_dim=1, end_dim=cond.ndim - 3)
|
||||
uncond = uncond.flatten(start_dim=1, end_dim=uncond.ndim - 3)
|
||||
return WCFGContext(
|
||||
cond=cond,
|
||||
uncond=uncond,
|
||||
x=x,
|
||||
sigma=sigma,
|
||||
wavelet=wavelet,
|
||||
dtype=eff_dtype,
|
||||
op_kwargs=op_kwargs,
|
||||
)
|
||||
|
||||
def process_output(
|
||||
self,
|
||||
*,
|
||||
result: torch.Tensor,
|
||||
rule: WCFGRule,
|
||||
ctx: WCFGContext,
|
||||
) -> torch.Tensor:
|
||||
x_shape = ctx.x.shape
|
||||
if rule.wavelet.use_1d_dwt:
|
||||
result = result[..., : ctx.cond.shape[2]].reshape(x_shape)
|
||||
elif ctx.x.ndim > 4:
|
||||
result = result[..., : x_shape[-2], : x_shape[-1]].reshape(x_shape)
|
||||
else:
|
||||
result = result[tuple(slice(None, sz) for sz in x_shape)]
|
||||
if rule.target_mode == WCFGTarget.DENOISED:
|
||||
result = ctx.x - result
|
||||
elif rule.target_mode == WCFGTarget.NOISE_NORM:
|
||||
result *= ctx.sigma
|
||||
return self.maybe_op(result, self.operation_wavelet_cfg, **ctx.op_kwargs)
|
||||
|
||||
@classmethod
|
||||
def wavelet_cfg(
|
||||
cls,
|
||||
*,
|
||||
rule: WCFGRule,
|
||||
ctx: WCFGContext,
|
||||
pcts: WCFGPercentages,
|
||||
) -> torch.Tensor:
|
||||
verbose = rule.verbose
|
||||
diff_blend_function = utils.BLENDING_MODES[rule.difference_blend_mode]
|
||||
condw = ctx.wavelet.forward(ctx.cond.to(dtype=ctx.dtype))
|
||||
uncondw = ctx.wavelet.forward(ctx.uncond.to(ctx.dtype))
|
||||
if rule.cond is not None:
|
||||
condw = rule.get_and_apply_scales("cond", pcts, *condw, verbose=verbose)
|
||||
if rule.uncond is not None:
|
||||
uncondw = rule.get_and_apply_scales(
|
||||
"uncond",
|
||||
pcts,
|
||||
*uncondw,
|
||||
verbose=verbose,
|
||||
)
|
||||
diffw = wavelet_blend(
|
||||
condw,
|
||||
uncondw,
|
||||
yl_factor=1.0,
|
||||
blend_function=lambda a, b, _t: a - b,
|
||||
)
|
||||
if rule.diff is not None:
|
||||
diffw = rule.get_and_apply_scales("diff", pcts, *diffw, verbose=verbose)
|
||||
resultw = wavelet_blend(
|
||||
uncondw,
|
||||
diffw,
|
||||
yl_factor=rule.difference_blend_strength.get_value(pcts),
|
||||
blend_function=diff_blend_function,
|
||||
)
|
||||
if rule.final is not None:
|
||||
resultw = rule.get_and_apply_scales(
|
||||
"final",
|
||||
pcts,
|
||||
*resultw,
|
||||
verbose=verbose,
|
||||
)
|
||||
return ctx.wavelet.inverse(*resultw).to(dtype=ctx.x.dtype)
|
||||
|
||||
def __call__(self, args: dict) -> torch.Tensor:
|
||||
sigma = args["sigma"]
|
||||
sigma_f = sigma.max().item()
|
||||
rule = self.rules.get_rule(sigma_f)
|
||||
if rule is None:
|
||||
return self.fallback_cfg_function(args)
|
||||
if rule.verbose:
|
||||
tqdm.write(
|
||||
f"\nWCFG: Rule matched, sigma={sigma_f:.4f}, rule={rule.pretty_non_default()}",
|
||||
)
|
||||
blend_function = utils.BLENDING_MODES[rule.blend_mode]
|
||||
model = args["model"]
|
||||
pcts = WCFGPercentages.build(
|
||||
ms=model.model_sampling,
|
||||
start_sigma=rule.start_sigma,
|
||||
end_sigma=rule.end_sigma,
|
||||
sigma=sigma_f,
|
||||
sigmas=args.get("model_options", {})
|
||||
.get("transformer_options", {})
|
||||
.get("sample_sigmas"),
|
||||
)
|
||||
wcfg_blend = rule.blend_strength.get_value(pcts)
|
||||
if rule.blend_mode == "lerp" and wcfg_blend == 0:
|
||||
return self.maybe_op(
|
||||
self.fallback_cfg_function(args),
|
||||
self.operation_fallback_cfg,
|
||||
sigma=sigma,
|
||||
cond=args["cond_denoised"],
|
||||
uncond=args["uncond_denoised"],
|
||||
raw_args=args,
|
||||
)
|
||||
ctx = self.get_context(rule=rule, args=args)
|
||||
result = self.wavelet_cfg(rule=rule, ctx=ctx, pcts=pcts)
|
||||
if rule.blend_mode != "lerp" or wcfg_blend != 1.0:
|
||||
normal_result = self.maybe_op(
|
||||
self.fallback_cfg_function(args),
|
||||
self.operation_fallback_cfg,
|
||||
**ctx.op_kwargs,
|
||||
)
|
||||
if rule.target_mode == WCFGTarget.DENOISED:
|
||||
normal_result = ctx.x - normal_result
|
||||
elif rule.target_mode == WCFGTarget.NOISE_NORM:
|
||||
normal_result /= ctx.sigma
|
||||
result = blend_function(normal_result, result, wcfg_blend)
|
||||
result = self.process_output(result=result, ctx=ctx, rule=rule)
|
||||
return self.maybe_op(
|
||||
result,
|
||||
self.operation_result,
|
||||
**ctx.op_kwargs,
|
||||
).contiguous()
|
||||
@@ -0,0 +1,240 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import torch
|
||||
|
||||
from .utils import fallback
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable, Sequence
|
||||
|
||||
try:
|
||||
import pytorch_wavelets as ptwav
|
||||
import pywt
|
||||
|
||||
HAVE_WAVELETS = True
|
||||
except ImportError:
|
||||
ptwav = None
|
||||
pywt = None
|
||||
HAVE_WAVELETS = False
|
||||
|
||||
|
||||
class Wavelet:
|
||||
DEFAULT_MODE = "symmetric"
|
||||
DEFAULT_LEVEL = 3
|
||||
DEFAULT_WAVE = "db4"
|
||||
DEFAULT_USE_1D_DWT = False
|
||||
DEFAULT_USE_DTCWT = False
|
||||
DEFAULT_QSHIFT = "qshift_a"
|
||||
DEFAULT_BIORT = "near_sym_a"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
wave: str = DEFAULT_WAVE,
|
||||
level: int = DEFAULT_LEVEL,
|
||||
mode: str = DEFAULT_MODE,
|
||||
use_1d_dwt: bool = DEFAULT_USE_1D_DWT,
|
||||
use_dtcwt: bool = DEFAULT_USE_DTCWT,
|
||||
biort: str = DEFAULT_BIORT,
|
||||
qshift: str = DEFAULT_QSHIFT,
|
||||
inv_wave: str | None = None,
|
||||
inv_mode: str | None = None,
|
||||
inv_biort: str | None = None,
|
||||
inv_qshift=None,
|
||||
device: str | torch.device | None = None,
|
||||
):
|
||||
if not HAVE_WAVELETS:
|
||||
raise RuntimeError(
|
||||
"Wavelet noise requires the pytorch_wavelets package to be installed in your Python environment",
|
||||
)
|
||||
inv_wave = fallback(inv_wave, wave)
|
||||
inv_mode = fallback(inv_mode, mode)
|
||||
inv_biort = fallback(inv_biort, biort)
|
||||
inv_qshift = fallback(inv_qshift, qshift)
|
||||
if use_dtcwt:
|
||||
fwdfun, invfun = ptwav.DTCWTForward, ptwav.DTCWTInverse
|
||||
elif use_1d_dwt:
|
||||
fwdfun, invfun = ptwav.DWT1DForward, ptwav.DWT1DInverse
|
||||
else:
|
||||
fwdfun, invfun = ptwav.DWTForward, ptwav.DWTInverse
|
||||
if use_dtcwt:
|
||||
self._wavelet_forward = fwdfun(
|
||||
J=level,
|
||||
mode=mode,
|
||||
biort=biort,
|
||||
qshift=qshift,
|
||||
)
|
||||
self._wavelet_inverse = invfun(
|
||||
mode=inv_mode,
|
||||
biort=inv_biort,
|
||||
qshift=inv_qshift,
|
||||
)
|
||||
else:
|
||||
self._wavelet_forward = fwdfun(J=level, wave=wave, mode=mode)
|
||||
self._wavelet_inverse = invfun(wave=inv_wave, mode=inv_mode)
|
||||
if device is not None:
|
||||
self._wavelet_forward = self._wavelet_forward.to(device=device)
|
||||
self._wavelet_inverse = self._wavelet_inverse.to(device=device)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
t: torch.Tensor,
|
||||
*,
|
||||
forward_function: Callable | None = None,
|
||||
) -> tuple[torch.Tensor, tuple]:
|
||||
return fallback(forward_function, self._wavelet_forward)(t)
|
||||
|
||||
def inverse(
|
||||
self,
|
||||
yl: torch.Tensor,
|
||||
yh: tuple,
|
||||
*,
|
||||
inverse_function: Callable | None = None,
|
||||
two_step_inverse: bool = False,
|
||||
) -> torch.Tensor:
|
||||
inverse_function = fallback(inverse_function, self._wavelet_inverse)
|
||||
if not two_step_inverse:
|
||||
return inverse_function((yl, yh))
|
||||
result = inverse_function((torch.zeros_like(yl), yh))
|
||||
result += inverse_function(
|
||||
(
|
||||
yl,
|
||||
tuple(torch.zeros_like(yh_band) for yh_band in yh),
|
||||
),
|
||||
)
|
||||
return result
|
||||
|
||||
def to(self, *args: Any, copy: bool = False, **kwargs: Any) -> Wavelet:
|
||||
o = Wavelet.__new__(Wavelet) if copy else self
|
||||
o._wavelet_forward = self._wavelet_forward.to(*args, **kwargs) # noqa: SLF001
|
||||
o._wavelet_inverse = self._wavelet_inverse.to(*args, **kwargs) # noqa: SLF001
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def wavelist() -> tuple:
|
||||
return tuple(pywt.wavelist()) if HAVE_WAVELETS else ()
|
||||
|
||||
@staticmethod
|
||||
def biortlist() -> tuple:
|
||||
return (
|
||||
("near_sym_a", "near_sym_b", "antonini", "legall") if HAVE_WAVELETS else ()
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def qshiftlist() -> tuple:
|
||||
return (
|
||||
("qshift_a", "qshift_b", "qshift_c", "qshift_d", "qshift_06")
|
||||
if HAVE_WAVELETS
|
||||
else ()
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def modelist() -> tuple:
|
||||
return (
|
||||
(
|
||||
"symmetric",
|
||||
"zero",
|
||||
"reflect",
|
||||
"replicate",
|
||||
"periodization",
|
||||
"periodic",
|
||||
"constant",
|
||||
)
|
||||
if HAVE_WAVELETS
|
||||
else ()
|
||||
)
|
||||
|
||||
|
||||
def expand_yh_scales(
|
||||
yh: Sequence,
|
||||
*,
|
||||
yh_scales: float | Sequence = 1.0,
|
||||
) -> float | tuple:
|
||||
yhlen = len(yh)
|
||||
yh_shape = yh[0].shape
|
||||
# Doesn't make sense to target orientations for 1D DWD (3D here).
|
||||
olen = yh_shape[2] if len(yh_shape) > 3 else 1
|
||||
# print(f"\nSIZES: yhlen={yhlen}, olen={olen}, yh_shape={yh[0].shape}")
|
||||
if isinstance(yh_scales, (float, int)):
|
||||
return ((float(yh_scales),) * olen,) * yhlen
|
||||
otemplate = (1.0,) * olen
|
||||
yh_scales = tuple(
|
||||
(float(band),) * olen
|
||||
if isinstance(band, (float, int))
|
||||
else (
|
||||
(
|
||||
*(float(i) for i in band[:olen]),
|
||||
*otemplate[: olen - len(band[:olen])],
|
||||
)
|
||||
if isinstance(band, (tuple, list))
|
||||
else band
|
||||
)
|
||||
for band in yh_scales
|
||||
)
|
||||
if "fill" in yh_scales:
|
||||
fillidx = yh_scales.index("fill")
|
||||
if "fill" in yh_scales[fillidx + 1 :]:
|
||||
raise ValueError("Only one fill allowed.")
|
||||
if fillidx == 0 or len(yh_scales) < 2:
|
||||
raise ValueError(
|
||||
"Invalid fill value, cannot be in the first position or the only item.",
|
||||
)
|
||||
yhslen = len(yh_scales)
|
||||
if yhslen - 1 < yhlen:
|
||||
# Need to pad.
|
||||
fill = (yh_scales[fillidx - 1],) * (yhlen - (len(yh_scales) - 1))
|
||||
yh_scales = (*yh_scales[:fillidx], *fill, *yh_scales[fillidx + 1 :])
|
||||
else:
|
||||
# Just remove the "fill".
|
||||
yh_scales = (*yh_scales[:fillidx], *yh_scales[fillidx + 1 :])
|
||||
return yh_scales[:yhlen]
|
||||
|
||||
|
||||
def wavelet_scaling(
|
||||
yl: torch.Tensor,
|
||||
yh: Sequence,
|
||||
yl_scale: float | torch.Tensor,
|
||||
yh_scales: float | Sequence | None,
|
||||
*,
|
||||
in_place: bool = False,
|
||||
) -> tuple:
|
||||
if not in_place:
|
||||
yl = yl.clone()
|
||||
yh = tuple(yhband.clone() for yhband in yh)
|
||||
if yl_scale != 1.0:
|
||||
yl *= yl_scale
|
||||
yh_scales = expand_yh_scales(
|
||||
yh,
|
||||
yh_scales=yh_scales if yh_scales is not None else 1.0,
|
||||
)
|
||||
for hscale, ht in zip(yh_scales, yh):
|
||||
if isinstance(hscale, (int, float)):
|
||||
ht *= hscale # noqa: PLW2901
|
||||
continue
|
||||
for lidx in range(min(ht.shape[2], len(hscale))):
|
||||
ht[:, :, lidx] *= hscale[lidx]
|
||||
return (yl, yh)
|
||||
|
||||
|
||||
def wavelet_blend(
|
||||
a: tuple,
|
||||
b: tuple,
|
||||
*,
|
||||
yl_factor: torch.Tensor | float,
|
||||
blend_function: Callable,
|
||||
yh_factor: torch.Tensor | float | None = None,
|
||||
yh_blend_function: Callable | None = None,
|
||||
) -> tuple:
|
||||
if not isinstance(yl_factor, torch.Tensor):
|
||||
yl_factor = a[0].new_full((1,), yl_factor)
|
||||
if yh_factor is None:
|
||||
yh_factor = yl_factor
|
||||
elif not isinstance(yh_factor, torch.Tensor):
|
||||
yh_factor = a[0].new_full((1,), yh_factor)
|
||||
yh_blend_function = fallback(yh_blend_function, blend_function)
|
||||
return (
|
||||
blend_function(a[0], b[0], yl_factor),
|
||||
tuple(yh_blend_function(ta, tb, yh_factor) for ta, tb in zip(a[1], b[1])),
|
||||
)
|
||||
@@ -7,6 +7,7 @@ ignore = [
|
||||
"ANN202",
|
||||
"ANN204",
|
||||
"ANN206",
|
||||
"ANN401",
|
||||
"C901",
|
||||
"CPY001",
|
||||
"DOC201",
|
||||
@@ -18,6 +19,7 @@ ignore = [
|
||||
"D105",
|
||||
"D106",
|
||||
"D107",
|
||||
"D401",
|
||||
"D211",
|
||||
"D213",
|
||||
"E402",
|
||||
@@ -29,10 +31,12 @@ ignore = [
|
||||
"FBT002",
|
||||
"PLR0912",
|
||||
"PLR0913",
|
||||
"PLR0914",
|
||||
"PLR0915",
|
||||
"PLR0917",
|
||||
"PLR2004",
|
||||
"T201",
|
||||
"TID252",
|
||||
"TRY003",
|
||||
"N802",
|
||||
"N999",
|
||||
|
||||
Reference in New Issue
Block a user