From 5280a2ca1174c721de457a1def5da3fb3281ecea Mon Sep 17 00:00:00 2001 From: blepping Date: Mon, 4 Aug 2025 14:42:43 -0600 Subject: [PATCH] Internal refactoring/cleanups Added a SonarAdvancedVoronoiNoise node --- README.md | 1 + changelog.md | 5 +- docs/advanced_noise_nodes.md | 50 ++++ docs/waveletcfg.md | 241 ++++++++++++++++++ py/nodes/__init__.py | 2 - py/nodes/misc.py | 238 ++++++++++++++++++ py/nodes/noise_filters.py | 13 +- py/nodes/noise_types.py | 128 +++++++++- py/noise.py | 103 +++++++- py/noise_generation.py | 462 ++++++++++++++++++++++++++++++++++ py/utils.py | 27 +- py/{nodes => }/wavelet_cfg.py | 379 ++++++---------------------- 12 files changed, 1338 insertions(+), 311 deletions(-) create mode 100644 docs/waveletcfg.md rename py/{nodes => }/wavelet_cfg.py (72%) diff --git a/README.md b/README.md index 90256de..9fd5e3b 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/changelog.md b/changelog.md index 8b7a6b7..327b462 100644 --- a/changelog.md +++ b/changelog.md @@ -2,13 +2,16 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. -## 20250727 +## 20250804 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 diff --git a/docs/advanced_noise_nodes.md b/docs/advanced_noise_nodes.md index 18f43e4..51104fa 100644 --- a/docs/advanced_noise_nodes.md +++ b/docs/advanced_noise_nodes.md @@ -448,3 +448,53 @@ 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` + +**Note**: Some modes act as wrappers for other modes. Unfortunately, there isn't currently a good way to escape parameters. Modes ignore parameters they don't understand and the wrapper modes will pass through any parameters they don't use themselves so you _can_ pass parameters to the submodes as long as they don't conflict (and only up to one level). + +#### 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` +* `manhatten` +* `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 result 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. + +#### 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. + +#### 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. diff --git a/docs/waveletcfg.md b/docs/waveletcfg.md new file mode 100644 index 0000000..2d66f10 --- /dev/null +++ b/docs/waveletcfg.md @@ -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). diff --git a/py/nodes/__init__.py b/py/nodes/__init__.py index 39ef10b..cd830e1 100644 --- a/py/nodes/__init__.py +++ b/py/nodes/__init__.py @@ -8,7 +8,6 @@ from . import ( noise_filters, noise_types, powernoise, - wavelet_cfg, ) NODE_CLASS_MAPPINGS = { @@ -26,7 +25,6 @@ for nm in ( noise_filters, noise_types, powernoise, - wavelet_cfg, ): NODE_CLASS_MAPPINGS |= getattr(nm, "NODE_CLASS_MAPPINGS", {}) NODE_DISPLAY_NAME_MAPPINGS |= getattr(nm, "NODE_DISPLAY_NAME_MAPPINGS", {}) diff --git a/py/nodes/misc.py b/py/nodes/misc.py index 9631dde..d05f040 100644 --- a/py/nodes/misc.py +++ b/py/nodes/misc.py @@ -10,10 +10,12 @@ import numpy as np import torch import yaml from comfy import model_management, samplers +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, @@ -659,10 +661,246 @@ class SonarSplitNoiseChainNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNode ) +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, } diff --git a/py/nodes/noise_filters.py b/py/nodes/noise_filters.py index 1eba841..91ab3a9 100644 --- a/py/nodes/noise_filters.py +++ b/py/nodes/noise_filters.py @@ -405,7 +405,7 @@ class SonarBlendedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix lambda: NoiseChainInputTypes() .req_float_noise_2_percent( default=0.5, - tooltip="Blend strength for custom_noise_2. Note that if set to 0 then custom_noise_2 is optional (and will not be called to generate noise) and if set to 1 then custom_noise_1 will not be called to generate noise. This is worth mentioning since going from a strength of 0.000000001 to 0 could make a big difference.", + tooltip="Blend strength for custom_noise_2. Note that if set to 0 then custom_noise_2 is optional (and will not be called to generate noise) and if set to 1 then custom_noise_1 will not be called to generate noise. This only applies when custom_noise_mask is not connected. This is worth mentioning since going from a strength of 0.000000001 to 0 could make a big difference. Important: When custom_noise_mask is connected, this value will be added to the mask and then the mask will be clamped to 0 through 1. In other words, you could use this to ensure the mask ranges between 0.5 and 1.0 by setting it to 0.5 or ensure it ranges between 0 and 0.5 by setting it to -0.5.", ) .req_selectblend( tooltip="Mode used for blending the two noise types. More modes will be available if ComfyUI-bleh is installed.", @@ -414,10 +414,13 @@ class SonarBlendedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix tooltip="Controls whether the generated noise is normalized to 1.0 strength. For weird blend modes, you may want to set this to forced.", ) .opt_customnoise_custom_noise_1( - tooltip="Custom noise. Optional if noise_2_percent is 1.", + tooltip="Custom noise. Optional if noise_2_percent is 1 and custom_noise_mask is not connected..", ) .opt_customnoise_custom_noise_2( - tooltip="Custom noise. Optional if noise_2_percent is 0.", + tooltip="Custom noise. Optional if noise_2_percent is 0 and custom_noise_mask is not connected..", + ) + .opt_customnoise_custom_noise_mask( + tooltip="Custom noise. If connected, this will be used instead of noise_2_percent to determine the blend ratio. Noise generated by this will be normalized to a 0 through 1 scale. When connected, both custom noise inputs are mandatory.", ), ) @@ -435,6 +438,7 @@ class SonarBlendedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix noise_2_percent, custom_noise_1=None, custom_noise_2=None, + custom_noise_mask=None, blend_mode="lerp", ): blend_function = utils.BLENDING_MODES.get(blend_mode) @@ -448,6 +452,7 @@ class SonarBlendedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix normalize=self.get_normalize(normalize), custom_noise_1=custom_noise_1, custom_noise_2=custom_noise_2, + custom_noise_mask=custom_noise_mask, noise_2_percent=noise_2_percent, ) @@ -1408,7 +1413,7 @@ class SonarCustomNoiseParametersNode( SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin, ): - DESCRIPTION = "Custom noise type that allows overriding some parameters." + DESCRIPTION = "Custom noise type that allows setting parameters like dtype or forking the RNG." INPUT_TYPES = SonarLazyInputTypes( lambda: NoiseNoChainInputTypes() diff --git a/py/nodes/noise_types.py b/py/nodes/noise_types.py index 8551822..b2f2f2d 100644 --- a/py/nodes/noise_types.py +++ b/py/nodes/noise_types.py @@ -3,7 +3,7 @@ from __future__ import annotations import torch from .. import noise, utils -from ..noise_generation import DistroNoiseGenerator +from ..noise_generation import DistroNoiseGenerator, VoronoiNoiseGenerator from .base import ( NoiseChainInputTypes, SonarCustomNoiseNodeBase, @@ -603,11 +603,137 @@ class SonarWaveletNoiseNode( ) +class SonarAdvancedVoronoiNoiseNode(SonarCustomNoiseNodeBase): + DESCRIPTION = "Voronoi noise is a very weird noise type. The default settings are just borderline usable with SDXL at a 20% ratio with normal Gaussian noise. I recommend reading the section on this noise type in the project documentation (under advanced noise types) as there are too many features to describe in the node itself." + + INPUT_TYPES = SonarLazyInputTypes( + lambda _pretty_distance_modes=", ".join( # noqa: B008 + sorted(VoronoiNoiseGenerator.voronoi_distance_modes), # noqa: B008 + ), + _pretty_result_modes=", ".join( # noqa: B008 + sorted(VoronoiNoiseGenerator.voronoi_result_modes), # noqa: B008 + ): NoiseChainInputTypes() + .req_string_n_points( + default="256", + tooltip="Controls the number of features points in the generated noise. Higher generally results in more detail/better results but is slower. May be a comma separated list for each octave (only applicable when octave mode is set to new_features). 2 is the minimum value.", + ) + .req_string_distance_mode( + default="euclidean", + placeholder=f"One of: {_pretty_distance_modes}", + tooltip="Distance modes. You can specify a comma-separated list of items which will be used for each octave.\n" + "You can specify an average of multiple distance modes by separating the names with +.\n" + "Some modes can take arguments. Example syntax: modename:argname=value:argname=value\n" + "All modes support scaling their output with dscale (which defaults to 1).\n" + f"Possible distance modes: {_pretty_distance_modes}", + ) + .req_float_z_initial( + default=0.0, + tooltip="Initial value for z (depth).", + ) + .req_float_z_increment( + default=1.0, + tooltip="Amount z (depth) is incremented when applicable.", + ) + .req_float_z_max( + default=9999.0, + tooltip="Maximum difference from the intial value. At that point, z_max_mode will apply. When set to 0, z_increment has no effect and you will get different noise each time you call the noise sampler.", + ) + .req_field_z_max_mode( + ( + "reset", + "wrap", + "bounce", + ), + default="reset", + tooltip="Controls what happens when the z_max limit is hit (see tooltip for z_max). Reset will reset the feature points and z to the initial values. Wrap will reset z to the initial value. Bounce will flip the sign on the increment and do an increment.", + ) + .req_string_result_mode( + default="diff2", + placeholder=f"One of: {_pretty_result_modes}", + tooltip="Result modes. You can specify a comma-separated list of items which will be used for each octave.\n" + "You can specify an average of multiple result modes by separating the names with +.\n" + "Some modes can take arguments. Example syntax: modename:argname=value:argname=value\n" + "All modes support scaling their output with rscale (which defaults to 1).\n" + f"Possible result modes: {_pretty_result_modes}", + ) + .req_field_octave_mode( + ("same_features", "new_features"), + default="new_features", + tooltip="Only relevant when generating multiple octaves. Controls whether octaves share a set of feature points or if they are different for each octave (note that this is slower).", + ) + .req_int_octaves( + default=3, + min=1, + tooltip="Number of octaves of noise to generate.", + ) + .req_float_gain(default=0.75) + .req_float_lacunarity(default=2.0) + .req_float_initial_amplitude(default=1.0) + .req_float_initial_scale(default=1.0) + .req_normalizetristate_normalize() + .opt_customnoise( + "custom_noise", + tooltip="Optional input if you want to use some other noise type for the initial feature points. Won't work well with noise types that care about the content of the latent (I think only spectral modulation) or manage their own seed (I believe this only applies to Brownian or if you're using the custom noise parameters node to override seeds/fork the RNG).", + ), + ) + + @classmethod + def get_item_class(cls): + return noise.AdvancedVoronoiNoise + + def go( + self, + *, + factor: float, + rescale: float, + n_points: str, + distance_mode: str, + z_initial: float, + z_increment: float, + z_max: float, + z_max_mode: str, + result_mode: str, + octave_mode: str, + octaves: int, + gain: float, + lacunarity: float, + initial_amplitude: float, + initial_scale: float, + normalize: str, + custom_noise=None, + sonar_custom_noise_opt=None, + ): + n_points = tuple(int(v) for v in n_points.split(",")) + distance_mode = tuple(v.strip() for v in distance_mode.split(",")) + result_mode = tuple(v.strip() for v in result_mode.split(",")) + return super().go( + factor, + rescale=rescale, + sonar_custom_noise_opt=sonar_custom_noise_opt, + n_points=n_points, + distance_mode=distance_mode, + z_initial=z_initial, + z_increment=z_increment, + z_max=z_max, + z_max_mode=z_max_mode, + result_mode=result_mode, + octave_mode=octave_mode, + octaves=octaves, + gain=gain, + lacunarity=lacunarity, + initial_amplitude=initial_amplitude, + initial_scale=initial_scale, + custom_noise=custom_noise, + normalize=normalize, + ) + + NODE_CLASS_MAPPINGS = { "SonarAdvancedPyramidNoise": SonarAdvancedPyramidNoiseNode, "SonarAdvanced1fNoise": SonarAdvanced1fNoiseNode, "SonarAdvancedPowerLawNoise": SonarAdvancedPowerLawNoiseNode, "SonarAdvancedCollatzNoise": SonarAdvancedCollatzNoiseNode, "SonarAdvancedDistroNoise": SonarAdvancedDistroNoiseNode, + "SonarAdvancedVoronoiNoise": SonarAdvancedVoronoiNoiseNode, "SonarWaveletNoise": SonarWaveletNoiseNode, } diff --git a/py/noise.py b/py/noise.py index 7b4a474..491b1b2 100644 --- a/py/noise.py +++ b/py/noise.py @@ -443,6 +443,30 @@ class AdvancedWaveletNoise(AdvancedNoiseBase): return result +class AdvancedVoronoiNoise(AdvancedNoiseBase): + ns_factory_arg_keys = tuple(VoronoiNoiseGenerator.ng_params(no_super=True)) + + @property + def ns_factory(self): + return VoronoiNoiseGenerator + + def clone_key(self, k): + if k == "custom_noise" and self.custom_noise is not None: + return self.custom_noise.clone() + return super().clone_key(k) + + def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + if x.ndim != 4: + raise ValueError("Can only handle 4+ dimensional latents") + return super().make_noise_sampler( + x, + *args, + normalized=normalized, + noise_sampler_factory=self.custom_noise, + **kwargs, + ) + + class CompositeNoise(CustomNoiseItemBase): def __init__( self, @@ -1284,17 +1308,26 @@ class BlendedNoise(CustomNoiseItemBase): blend_function, custom_noise_1=None, custom_noise_2=None, + custom_noise_mask=None, noise_2_percent=0.5, ): - if custom_noise_1 is None and noise_2_percent != 1: + if custom_noise_1 is None and ( + custom_noise_mask is not None or noise_2_percent != 1 + ): raise ValueError( "When custom_noise_1 is not attached noise_2_percent must be set to 1", ) - if custom_noise_2 is None and noise_2_percent != 0: + if custom_noise_2 is None and ( + custom_noise_mask is not None or noise_2_percent != 0 + ): raise ValueError( "When custom_noise_2 is not attached noise_2_percent must be set to 0", ) - if noise_2_percent == 1 and custom_noise_1 is None: + if ( + custom_noise_mask is None + and noise_2_percent == 1 + and custom_noise_1 is None + ): custom_noise_1, custom_noise_2 = custom_noise_2, None noise_2_percent = 0.0 super().__init__( @@ -1303,6 +1336,9 @@ class BlendedNoise(CustomNoiseItemBase): blend_function=blend_function, custom_noise_1=custom_noise_1.clone(), custom_noise_2=None if custom_noise_2 is None else custom_noise_2.clone(), + custom_noise_mask=None + if custom_noise_mask is None + else custom_noise_mask.clone(), normalize=normalize, ) @@ -1311,6 +1347,12 @@ class BlendedNoise(CustomNoiseItemBase): return self.custom_noise_1.clone() if k == "custom_noise_2": return None if self.custom_noise_2 is None else self.custom_noise_2.clone() + if k == "custom_noise_mask": + return ( + None + if self.custom_noise_mask is None + else self.custom_noise_mask.clone() + ) return super().clone_key(k) def make_noise_sampler(self, x, *args, normalized=True, **kwargs): @@ -1318,7 +1360,7 @@ class BlendedNoise(CustomNoiseItemBase): normalize = self.get_normalize("normalize", normalized) blend_function = self.blend_function n2_blend = self.noise_2_percent - n2_blend_tensor = x.new_full((1,), n2_blend) + ns_1 = self.custom_noise_1.make_noise_sampler( x, *args, @@ -1335,10 +1377,26 @@ class BlendedNoise(CustomNoiseItemBase): **kwargs, ) ) + ns_mask = ( + None + if self.custom_noise_mask is None + else self.custom_noise_mask.make_noise_sampler( + x, + *args, + normalized=False, + **kwargs, + ) + ) + n2_blend_tensor = x.new_full((1,), n2_blend) if ns_mask is None else None def noise_sampler(s, sn): + nonlocal n2_blend_tensor noise_1 = ns_1(s, sn) noise_2 = None if ns_2 is None else ns_2(s, sn) + if ns_mask is not None: + n2_blend_tensor = ( + utils.normalize_to_scale(ns_mask(s, sn), 0.0, 1.0) + n2_blend + ).clamp_(0.0, 1.0) noise = ( noise_1 if noise_2 is None @@ -2311,6 +2369,43 @@ NOISE_SAMPLERS: dict[NoiseType, Callable] = { ), ), NoiseType.COLLATZ: NoiseSampler.wrap(CollatzNoiseGenerator), + NoiseType.VORONOI_FUZZ: NoiseSampler.wrap( + partial( + VoronoiNoiseGenerator, + n_points=(256,), + octaves=1, + distance_mode=("fuzz:name=angle_tanh:fuzz=0.1",), + result_mode=("diff2",), + z_max=0.0, + ), + ), + NoiseType.VORONOI_MIX: NoiseSampler.wrap( + partial( + MixedNoiseGenerator, + name="voronoi_mix", + noise_mix=( + ( + VoronoiNoiseGenerator, + { + "n_points": (256,), + "octaves": 3, + "distance_mode": ("euclidean",), + "result_mode": ("diff2",), + "octave_mode": "new_features", + "lacunarity": 2.0, + "gain": 0.75, + "z_max": 0.0, + }, + lambda t: t.mul_(0.6), + ), + ( + GaussianNoiseGenerator, + {}, + lambda t: t.mul_(0.4), + ), + ), + ), + ), } diff --git a/py/noise_generation.py b/py/noise_generation.py index 71225bc..05d774c 100644 --- a/py/noise_generation.py +++ b/py/noise_generation.py @@ -62,6 +62,8 @@ class NoiseType(Enum): UNIFORM = auto() VELVET = auto() VIOLET = auto() + VORONOI_FUZZ = auto() + VORONOI_MIX = auto() WAVELET = auto() WHITE = auto() @@ -1284,6 +1286,465 @@ class PowerOldNoiseGenerator(NoiseGenerator): return noise.sub_(mean).div_(std) +# With help from ChatGPT. +class VoronoiNoiseGenerator(NoiseGenerator): + name = "voronoi" + MIN_DIMS = 4 + MAX_DIMS = 4 + + @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: + return self.feature_points[octave % len(self.feature_points)] + + 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)] + + voronoi_distance_modes = frozenset(( + "euclidean", + "manhatten", + "chebyshev", + "minkowski", + "quadratic", + "angle", + "angle_tanh", + "angle_sigmoid", + "fuzz", + )) + + @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_fuzz( + self, + *args, + name: str = "f1", + fuzz: float | str = 0.25, + **kwargs, + ) -> torch.Tensor: + fuzz = float(fuzz) + name = name.strip().lower() + if name not in self.voronoi_distance_modes: + errstr = f"Bad voronoi fuzz distance mode name: {name}" + raise ValueError(errstr) + result = getattr(self, f"_voronoi_distance_{name}")(*args, **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(v.strip() for v 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 + if mode_name not in self.voronoi_distance_modes: + errstr = f"Bad distance mode {mode}" + raise ValueError(errstr) + handler = getattr(self, f"_voronoi_distance_{mode_name}") + curr_result = handler(d, **mode_kwargs).mul_(result_scale) + result = curr_result if result is None else result.add_(curr_result) + return result + + voronoi_result_modes = frozenset(( + "f", + "f1", + "f2", + "f3", + "f4", + "diff", + "diff2", + "inv_f", + "inv_f1", + "inv_f2", + "inv_f3", + "inv_f4", + "cellid", + "ridge", + "median_distance", + "fuzz", + )) + + @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: + name = name.strip().lower() + if name not in self.voronoi_result_modes: + errstr = f"Bad voronoi ridge result mode name: {name}" + raise ValueError(errstr) + return 1.0 - ( + float(exp) * getattr(self, f"_voronoi_result_{name}")(*args, **kwargs) + ) + + @staticmethod + def _voronoi_result_median_distance( + *_args, + get_sorted: Callable, + **_kwargs, + ) -> torch.Tensor: + return get_sorted().median(dim=-1).values + + def _voronoi_result_fuzz( + self, + *args, + name: str = "f1", + fuzz: float | str = 0.25, + **kwargs, + ) -> torch.Tensor: + fuzz = float(fuzz) + name = name.strip().lower() + if name not in self.voronoi_result_modes: + errstr = f"Bad voronoi fuzz result mode name: {name}" + raise ValueError(errstr) + result = getattr(self, f"_voronoi_result_{name}")(*args, **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_result(self, d: 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 + + 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 = {} + if mode_name not in self.voronoi_result_modes: + errstr = f"Bad result mode {mode}" + raise ValueError(errstr) + handler = getattr(self, f"_voronoi_result_{mode_name}") + curr_result = handler( + d, + get_sorted=get_sorted, + **mode_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 = (grid_3d - fp + 0.5) % 1.0 - 0.5 # Wrap to [-0.5, 0.5) + + d = self.voronoi_distance(d, octave=octave) + return self.voronoi_result(d, 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 + + # Idea from https://github.com/ClownsharkBatwing/RES4LYF/ (wave and mode defaults also from that source) class WaveletFilteredNoiseGenerator(FramesToChannelsNoiseGenerator): name = "waveletfilter" @@ -2016,6 +2477,7 @@ __all__ = ( "ScatternetFilteredNoiseGenerator", "StudentTNoiseGenerator", "UniformNoiseGenerator", + "VoronoiNoiseGenerator", "WaveletFilteredNoiseGenerator", "WaveletNoiseGenerator", ) diff --git a/py/utils.py b/py/utils.py index c7b72af..89f9b51 100644 --- a/py/utils.py +++ b/py/utils.py @@ -3,7 +3,7 @@ from __future__ import annotations import math import random from functools import partial -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Callable import torch from comfy.model_management import device_supports_non_blocking, get_torch_device @@ -30,6 +30,31 @@ UPSCALE_METHODS = ( ) +def blend_scalar( + a: float, + b: float, + t: float, + *, + blend_function: Callable | None = None, + clamp_function: Callable | None = None, +) -> float: + if blend_function is None: + return maybe_apply( + a * (1.0 - t) + b * t, + clamp_function is not None, + clamp_function, + ) + return maybe_apply( + blend_function( + *(torch.tensor((v,), device="cpu", dtype=torch.float64) for v in (a, b, t)), + ) + .cpu() + .item(), + clamp_function is not None, + clamp_function, + ) + + def scale_samples( samples: torch.Tensor, width: int, diff --git a/py/nodes/wavelet_cfg.py b/py/wavelet_cfg.py similarity index 72% rename from py/nodes/wavelet_cfg.py rename to py/wavelet_cfg.py index 09d786d..da854bb 100644 --- a/py/nodes/wavelet_cfg.py +++ b/py/wavelet_cfg.py @@ -5,23 +5,31 @@ from enum import Enum, auto from typing import TYPE_CHECKING, Callable, NamedTuple import torch -import yaml from tqdm import tqdm -from .. import utils -from ..external import IntegratedNode -from ..wavelet_functions import ( +from . import utils +from .wavelet_functions import ( Wavelet, expand_yh_scales, wavelet_blend, wavelet_scaling, ) -from .base import SonarInputTypes, SonarLazyInputTypes 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() @@ -304,14 +312,7 @@ class WCFGScheduledScale(NamedTuple): return pct def pretty_non_default(self) -> 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(self, _fn)) for _fn in self._fields) - if fv != getattr(DEFAULT_SCHEDULEDSCALE, fn) - ) - return f"WCFGScheduledScale({result})" + return pretty_non_default(self, defaults=DEFAULT_SCHEDULEDSCALE) DEFAULT_SCHEDULEDSCALE = WCFGScheduledScale() @@ -321,6 +322,7 @@ 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: @@ -336,6 +338,7 @@ class WCFGScalesRange(NamedTuple): 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), @@ -344,6 +347,7 @@ class WCFGScalesRange(NamedTuple): bool(scales_end), WCFGScheduledScale.build, ), + blend_mode=blend_mode, ) def get_scales( @@ -359,9 +363,10 @@ class WCFGScalesRange(NamedTuple): if verbose: tqdm.write(f"WCFG: pct={pct:.4f}, percentages: {pcts}") start, end = self.scales_start, self.scales_end - if pct <= 0: + simple_blend = self.blend_mode == "lerp" + if pct <= 0 and simple_blend: simple_result = start - elif pct >= 1: + elif pct >= 1 and simple_blend: simple_result = end else: simple_result = None @@ -371,12 +376,22 @@ class WCFGScalesRange(NamedTuple): f"WCFG: {simple_result.pretty_scales()}", ) return simple_result - start_scale, end_scale = 1.0 - pct, pct start_yh_scales = expand_yh_scales(yh, yh_scales=start.yh_scales) end_yh_scales = expand_yh_scales(yh, yh_scales=end.yh_scales) - yl_scale = start.yl_scale * start_scale + end.yl_scale * end_scale + 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(os * start_scale + oe * end_scale for os, oe in zip(bs, be)) + 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) @@ -404,14 +419,7 @@ class WCFGScalesRange(NamedTuple): return self.get_scales(pcts, yh, verbose=verbose).apply_scales(yl, yh) def pretty_non_default(self) -> 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(self, _fn)) for _fn in self._fields) - if fv != getattr(DEFAULT_SCALESRANGE, fn) - ) - return f"WCFGScalesRange({result})" + return pretty_non_default(self, defaults=DEFAULT_SCALESRANGE) DEFAULT_SCALESRANGE = WCFGScalesRange() @@ -457,6 +465,46 @@ class WCFGScheduledFloat(NamedTuple): 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 @@ -469,18 +517,8 @@ class WCFGRule(NamedTuple): cond: WCFGScalesRange | WCFGScales | None = None uncond: WCFGScalesRange | WCFGScales | None = None final: WCFGScalesRange | WCFGScales | None = None - 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" + wavelet: WCFGWaveletSettings = DEFAULT_WAVELETSETTINGS high_precision_mode: bool = True - inv_wave: str | None = None - inv_padding_mode: str | None = None - inv_biort: str | None = None - inv_qshift: str | None = None difference_blend_mode: str = "inject" difference_blend_strength: WCFGScheduledFloat = WCFGScheduledFloat(1.0) @@ -519,24 +557,12 @@ class WCFGRule(NamedTuple): 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 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, - ) + return self.wavelet.make_wavelet(**kwargs) def get_and_apply_scales( self, @@ -555,14 +581,7 @@ class WCFGRule(NamedTuple): return scales.apply_scales(yl, yh) def pretty_non_default(self) -> 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(self, _fn)) for _fn in self._fields) - if fv != getattr(DEFAULT_RULE, fn) - ) - return f"WCFGRule({result})" + return pretty_non_default(self, defaults=DEFAULT_RULE) DEFAULT_RULE = WCFGRule() @@ -659,7 +678,7 @@ class WaveletCFG: sigma_orig = sigma = args["sigma"] rule_id = id(rule) x = args["input"] - if x.ndim == 3 and not rule.use_1d_dwt: + 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( @@ -691,7 +710,7 @@ class WaveletCFG: wavelet = rule.make_wavelet() self.wavelet_cache[rule_id] = wavelet wavelet = wavelet.to(device=x.device, dtype=eff_dtype) - if rule.use_1d_dwt: + if rule.wavelet.use_1d_dwt: cond = cond.flatten(start_dim=2) uncond = uncond.flatten(start_dim=2) elif x.ndim > 4: @@ -715,7 +734,7 @@ class WaveletCFG: ctx: WCFGContext, ) -> torch.Tensor: x_shape = ctx.x.shape - if rule.use_1d_dwt: + 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) @@ -821,239 +840,3 @@ class WaveletCFG: self.operation_result, **ctx.op_kwargs, ).contiguous() - - -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. -# Note: Do not remove keys and there isn't really any error checking. -# 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: 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 - - # 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 - - -# 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 not 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 = { - "SonarWaveletCFG": SonarWaveletCFGNode, -}