From cf90ae74e10b70612219754e4a3cf781306a1405 Mon Sep 17 00:00:00 2001 From: blepping <157360029+blepping@users.noreply.github.com> Date: Tue, 5 Aug 2025 17:07:29 -0600 Subject: [PATCH] Input types refactor and wavelet CFG (#18) * 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. --- README.md | 1 + __init__.py | 14 +- changelog.md | 11 + docs/advanced_noise_nodes.md | 50 + docs/base_noise_types.md | 2 + docs/waveletcfg.md | 241 ++++ py/external.py | 6 +- py/latent_ops.py | 50 +- py/nodes/__init__.py | 33 +- py/nodes/base.py | 271 +++-- py/nodes/base_inputtypes.py | 263 +++++ py/{ => nodes}/freeu_extreme.py | 279 ++--- py/nodes/integrations.py | 255 ++-- py/nodes/latent_operations.py | 519 ++++----- py/nodes/misc.py | 825 ++++++------- py/nodes/momentum_samplers.py | 282 ++--- py/nodes/noise_filters.py | 1920 ++++++++++++++----------------- py/nodes/noise_types.py | 853 ++++++-------- py/{ => nodes}/powernoise.py | 286 ++--- py/noise.py | 269 ++++- py/noise_generation.py | 490 +++++++- py/utils.py | 184 ++- py/wavelet_cfg.py | 842 ++++++++++++++ py/wavelet_functions.py | 108 +- ruff.toml | 2 + 25 files changed, 4932 insertions(+), 3124 deletions(-) create mode 100644 docs/waveletcfg.md create mode 100644 py/nodes/base_inputtypes.py rename py/{ => nodes}/freeu_extreme.py (50%) rename py/{ => nodes}/powernoise.py (79%) create mode 100644 py/wavelet_cfg.py 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/__init__.py b/__init__.py index a5c338f..bf5d568 100644 --- a/__init__.py +++ b/__init__.py @@ -1,7 +1,7 @@ import sys from . import py # noqa: F401 -from .py import freeu_extreme, nodes, powernoise, sonar +from .py import nodes, sonar def blep_init(): @@ -15,17 +15,9 @@ def blep_init(): sonar.add_samplers() blep_init() -NODE_CLASS_MAPPINGS = ( - nodes.NODE_CLASS_MAPPINGS - | powernoise.NODE_CLASS_MAPPINGS - | freeu_extreme.NODE_CLASS_MAPPINGS -) +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", {}) - | getattr(powernoise, "NODE_DISPLAY_NAME_MAPPINGS", {}) - | getattr(freeu_extreme, "NODE_DISPLAY_NAME_MAPPINGS", {}) -) +NODE_DISPLAY_NAME_MAPPINGS = getattr(nodes, "NODE_DISPLAY_NAME_MAPPINGS", {}) __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/changelog.md b/changelog.md index c074c2c..e5b760f 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,17 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. +## 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. 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/base_noise_types.md b/docs/base_noise_types.md index 92ecd18..cb6c258 100644 --- a/docs/base_noise_types.md +++ b/docs/base_noise_types.md @@ -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 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/external.py b/py/external.py index 09b0ede..aa01139 100644 --- a/py/external.py +++ b/py/external.py @@ -120,7 +120,11 @@ class IntegratedNode(type): 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 diff --git a/py/latent_ops.py b/py/latent_ops.py index f29bcec..53553ed 100644 --- a/py/latent_ops.py +++ b/py/latent_ops.py @@ -2,14 +2,18 @@ from __future__ import annotations import math import random +from typing import TYPE_CHECKING import torch from . import utils +if TYPE_CHECKING: + from types import Sequence + class SonarLatentOperation: - SKIP_ARGS = frozenset(("sigma", "t2", "cond", "uncond", "cond_scale", "raw_args")) + EXTENDED_LATENT_OPERATION = True def __init__( self, @@ -23,6 +27,8 @@ class SonarLatentOperation: 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( @@ -36,9 +42,9 @@ class SonarLatentOperation: op = self.op if op is None: return t - if not isinstance(op, SonarLatentOperation): - kwargs = {k: v for k, v in kwargs.items() if k not in self.SKIP_ARGS} - return op(t, *args, **kwargs) + if not getattr(op, "EXTENDED_LATENT_OPERATION", False): + return op(latent=t) + return op(*args, latent=t, **kwargs) def __call__( self, @@ -61,6 +67,7 @@ class SonarLatentOperationAdvanced(SonarLatentOperation): input_multiplier: float, output_multiplier: float, difference_multiplier: float, + ops: Sequence, op_alt=None, **kwargs: dict, ) -> None: @@ -71,6 +78,7 @@ class SonarLatentOperationAdvanced(SonarLatentOperation): self.output_multiplier = output_multiplier self.difference_multiplier = difference_multiplier self.op_alt = op_alt + self.ops = ops def __call__( self, @@ -87,14 +95,12 @@ class SonarLatentOperationAdvanced(SonarLatentOperation): if self.op_alt is None else self.call_op(t, sigma=sigma, op=self.op_alt, **kwargs) ) - output = self.call_op( - t if self.input_multiplier == 1.0 else t * self.input_multiplier, - sigma=sigma, - **kwargs, - ) - if self.output_multiplier != 1.0: - output = output * self.output_multiplier # noqa: PLR6104 - diff = output - t + 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 return self.blend_function(t, diff, self.blend_strength) @@ -181,11 +187,23 @@ class SonarLatentOperationNoise(SonarLatentOperation): class SonarLatentOperationSetSeed(SonarLatentOperation): - def __init__(self, *args: list, seed: int, **kwargs: dict): + def __init__(self, *args: list, seed: int, restore_rng_state: bool, **kwargs: dict): super().__init__(*args, **kwargs) self.seed = seed + self.restore_rng_state = restore_rng_state def __call__(self, *args: list, **kwargs: dict) -> torch.Tensor: - torch.manual_seed(self.seed) - random.seed(self.seed) - return super().__call__(*args, **kwargs) + 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 diff --git a/py/nodes/__init__.py b/py/nodes/__init__.py index 7289d16..cd830e1 100644 --- a/py/nodes/__init__.py +++ b/py/nodes/__init__.py @@ -1,31 +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, -} | ( - integrations.NODE_CLASS_MAPPINGS - | latent_operations.NODE_CLASS_MAPPINGS - | misc.NODE_CLASS_MAPPINGS - | momentum_samplers.NODE_CLASS_MAPPINGS - | noise_filters.NODE_CLASS_MAPPINGS - | noise_types.NODE_CLASS_MAPPINGS -) +} +NODE_DISPLAY_NAME_MAPPINGS = {} - -NODE_DISPLAY_NAME_MAPPINGS = ( - getattr(integrations, "NODE_DISPLAY_NAME_MAPPINGS", {}) - | getattr(latent_operations, "NODE_DISPLAY_NAME_MAPPINGS", {}) - | getattr(misc, "NODE_DISPLAY_NAME_MAPPINGS", {}) - | getattr(momentum_samplers, "NODE_DISPLAY_NAME_MAPPINGS", {}) - | getattr(noise_filters, "NODE_DISPLAY_NAME_MAPPINGS", {}) - | getattr(noise_types, "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", {}) diff --git a/py/nodes/base.py b/py/nodes/base.py index 566971c..d13fa87 100644 --- a/py/nodes/base.py +++ b/py/nodes/base.py @@ -1,11 +1,11 @@ -# ruff: noqa: TID252 from __future__ import annotations import abc from typing import Any -from .. import noise -from ..external import IntegratedNode +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 @@ -47,6 +47,151 @@ NOISE_INPUT_TYPES_HINT = ( ) +class SonarInputCollection(InputCollection): + def __init__(self, *args: list, **kwargs: dict): + super().__init__(*args, **kwargs) + self._DELEGATE_KEYS = self._DELEGATE_KEYS | frozenset(( # noqa: PLR6104 + "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: dict, + ): + 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: dict, + ) -> 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: dict, + ) -> 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: dict, + ) -> 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: dict, + ) -> 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: dict, + ): + return self.field( + name, + ("default", "forced", "disabled"), + default=default, + tooltip=tooltip, + **kwargs, + ) + + def floatpct(self, name: str, *, min=0.0, max=1.0, **kwargs: dict): # noqa: A002 + return self.float(name=name, min=min, max=max, **kwargs) + + +class SonarInputTypes(InputTypes): + _NO_REPLACE = True + + def __init__(self, *args: list, **kwargs: dict): + super().__init__( + *args, + collection_class=SonarInputCollection, + **kwargs, + ) + + +class SonarLazyInputTypes(LazyInputTypes): + _NO_REPLACE = True + + def __init__(self, *args: list, initializers=(MODULES.initialize,), **kwargs: dict): + super().__init__( + *args, + initializers=initializers, + **kwargs, + ) + + class SonarCustomNoiseNodeBase(metaclass=IntegratedNode): DESCRIPTION = "A custom noise item." RETURN_TYPES = ("SONAR_CUSTOM_NOISE",) @@ -58,48 +203,24 @@ class SonarCustomNoiseNodeBase(metaclass=IntegratedNode): def get_item_class(self): raise NotImplementedError - @classmethod - def INPUT_TYPES(cls, *, include_rescale=True, include_chain=True): - result = { - "required": { - "factor": ( - "FLOAT", - { - "default": 1.0, - "min": -10000.0, - "max": 10000.0, - "step": 0.001, - "round": False, - "tooltip": "Scaling factor for the generated noise of this type.", - }, - ), - }, - "optional": {}, - } - if include_rescale: - result["required"] |= { - "rescale": ( - "FLOAT", - { - "default": 0.0, - "min": 0.0, - "max": 10000.0, - "step": 0.001, - "round": False, - "tooltip": "When non-zero, this custom noise item and other custom noise items items connected to it will have their factor scaled to add up to the specified rescale value. When set to 0, rescaling is disabled.", - }, - ), - } - if include_chain: - result["optional"] |= { - "sonar_custom_noise_opt": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional input for more custom noise items.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - } - return result + 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, @@ -118,19 +239,35 @@ class SonarCustomNoiseNodeBase(metaclass=IntegratedNode): return (nis if rescale == 0 else nis.rescaled(rescale),) +class NoiseChainInputTypes(SonarInputTypes): + def __init__(self, *, parent=SonarCustomNoiseNodeBase, **kwargs: dict): + super().__init__(parent=parent, **kwargs) + + +class NoiseNoChainInputTypes(SonarInputTypes): + def __init__( + self, + *, + parent=SonarCustomNoiseNodeBase, + parent_args=(), + parent_kwargs=None, + **kwargs: dict, + ): + 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): - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "noise_type": ( - tuple(noise.NoiseType.get_names()), - { - "tooltip": "Sets the type of noise to generate.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes().req_selectnoise_noise_type( + tooltip="Sets the type of noise to generate.", + ), + ) @classmethod def get_item_class(cls): @@ -140,21 +277,11 @@ class SonarCustomNoiseNode(SonarCustomNoiseNodeBase): class SonarCustomNoiseAdvNode(SonarCustomNoiseNode): DESCRIPTION = "A custom noise item allowing advanced YAML parameter input." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["optional"] |= { - "yaml_parameters": ( - "STRING", - { - "tooltip": "Allows specifying custom parameters via YAML. Note: When specifying paramaters this way, there is no error checking.", - "placeholder": "# YAML or JSON here", - "dynamicPrompts": False, - "multiline": True, - }, - ), - } - return result + 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: diff --git a/py/nodes/base_inputtypes.py b/py/nodes/base_inputtypes.py new file mode 100644 index 0000000..1f3d9b5 --- /dev/null +++ b/py/nodes/base_inputtypes.py @@ -0,0 +1,263 @@ +# ruff: noqa: A002 +from __future__ import annotations + +from copy import deepcopy +from functools import partial +from typing import Callable, TypeVar + + +class InputCollection: + _DELEGATE_KEYS = frozenset(( + "bool", + "boolean", + "clip", + "conditioning", + "field", + "float", + "image", + "int", + "latent", + "model", + "sampler", + "seed", + "sigmas", + "string", + "vae", + )) + + def __init__(self, **kwargs: dict): + 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) -> int: + return len(self.fields) + + def __contains__(self, key: str) -> bool: + return key in self.fields + + def field( + self, + name: str, + type: str | tuple, + *, + _skip: bool = False, + **kwargs: dict, + ) -> InputCollection: + if not _skip: + self.fields[name] = (type,) if not kwargs else (type, kwargs) + return self + + def string( + self, + name: str, + **kwargs: dict, + ) -> InputCollection: + return self.field(name, "STRING", **kwargs) + + def float( + self, + name: str, + *, + step: float = 0.001, + min: float = -10000.0, + max: float = 10000.0, + round: bool = False, + **kwargs: dict, + ) -> InputCollection: + return self.field( + name, + "FLOAT", + step=step, + min=min, + max=max, + round=round, + **kwargs, + ) + + def int( + self, + name: str, + *, + min: float = -10000, + max: float = 10000, + **kwargs: dict, + ) -> InputCollection: + return self.field( + name, + "INT", + min=min, + max=max, + **kwargs, + ) + + def bool( + self, + name: str, + default: bool = False, + **kwargs: dict, + ) -> InputCollection: + return self.field(name, "BOOLEAN", default=default, **kwargs) + + boolean = bool + + def seed( + self, + name: str = "seed", + *, + default: int = 0, + min: int = 0, + max: int = 0xFFFFFFFFFFFFFFFF, + tooltip="Seed to use for generated noise", + **kwargs: dict, + ) -> InputCollection: + return self.int( + name, + default=default, + min=min, + max=max, + tooltip=tooltip, + **kwargs, + ) + + def image(self, name: str = "image", **kwargs: dict) -> InputCollection: + return self.field(name, "IMAGE", **kwargs) + + def latent(self, name: str = "latent", **kwargs: dict) -> InputCollection: + return self.field(name, "LATENT", **kwargs) + + def conditioning( + self, + name: str = "conditioning", + **kwargs: dict, + ) -> InputCollection: + return self.field(name, "CONDITIONING", **kwargs) + + def model(self, name: str = "model", **kwargs: dict) -> InputCollection: + return self.field(name, "MODEL", **kwargs) + + def sigmas(self, name: str = "sigmas", **kwargs: dict) -> InputCollection: + return self.field(name, "SIGMAS", **kwargs) + + def sampler(self, name: str = "sampler", **kwargs: dict) -> InputCollection: + return self.field(name, "SAMPLER", **kwargs) + + def clip(self, name: str = "clip", **kwargs: dict) -> InputCollection: + return self.field(name, "CLIP", **kwargs) + + def vae(self, name: str = "vae", **kwargs: dict) -> 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: list, **kwargs: dict): + 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: list, **kwargs: dict): + 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: list, **kwargs: dict) -> dict: + return self.get_input_types(*args, **kwargs)() diff --git a/py/freeu_extreme.py b/py/nodes/freeu_extreme.py similarity index 50% rename from py/freeu_extreme.py rename to py/nodes/freeu_extreme.py index 319bf1b..2e6255d 100644 --- a/py/freeu_extreme.py +++ b/py/nodes/freeu_extreme.py @@ -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"" -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( diff --git a/py/nodes/integrations.py b/py/nodes/integrations.py index 94dc08a..a42af7a 100644 --- a/py/nodes/integrations.py +++ b/py/nodes/integrations.py @@ -1,15 +1,13 @@ -# ruff: noqa: TID252 - from __future__ import annotations from comfy import samplers -from .. import external, noise, utils +from .. import external, noise from .base import ( - NOISE_INPUT_TYPES_HINT, - WILDCARD_NOISE, - IntegratedNode, + NoiseNoChainInputTypes, SonarCustomNoiseNodeBase, + SonarInputTypes, + SonarLazyInputTypes, SonarNormalizeNoiseNodeMixin, ) @@ -25,68 +23,28 @@ class SonarBlendFilterNoiseNode( ): DESCRIPTION = "Custom noise type that allows blending and filtering the output of another noise generator using ComfyUI-bleh." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - bleh_filter_presets = ( - () if bleh is None else tuple(bleh.py.latent_utils.FILTER_PRESETS.keys()) + 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()), ) - bleh_enhance_methods = ( - () if bleh is None else ("none", *bleh.py.latent_utils.ENHANCE_METHODS) + .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", ) - result["required"] |= { - "sonar_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "blend_mode": ( - ("simple_add", *utils.BLENDING_MODES.keys()), - {"default": "simple_add"}, - ), - "ffilter": (bleh_filter_presets,), - "ffilter_custom": ("STRING", {"default": ""}), - "ffilter_scale": ( - "FLOAT", - { - "default": 1.0, - "min": -100.0, - "max": 100.0, - "step": 0.001, - "round": False, - }, - ), - "ffilter_strength": ( - "FLOAT", - { - "default": 0.0, - "min": -100.0, - "max": 100.0, - "step": 0.001, - "round": False, - }, - ), - "ffilter_threshold": ( - "INT", - {"default": 1, "min": 1, "max": 32}, - ), - "enhance_mode": (bleh_enhance_methods,), - "enhance_strength": ( - "FLOAT", - { - "default": 0.0, - "min": -100.0, - "max": 100.0, - "step": 0.001, - "round": False, - }, - ), - "affect": (("result", "noise", "both"),), - "normalize_result": (("default", "forced", "disabled"),), - "normalize_noise": (("default", "forced", "disabled"),), - } - return result + .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): @@ -122,6 +80,8 @@ class SonarBlendFilterNoiseNode( ) 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( @@ -148,33 +108,15 @@ class SonarBlehOpsNoiseNode( "Custom noise type that allows manipulating noise with ComfyUI-bleh ops." ) - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - result["required"] |= { - "sonar_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - "rules": ( - "STRING", - { - "tooltip": "Enter rules in the bleh block ops format here.", - "placeholder": "# YAML ops here", - "dynamicPrompts": False, - "multiline": True, - }, - ), - } - return result + 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): @@ -201,69 +143,48 @@ class SonarBlehOpsNoiseNode( restart = None -class KRestartSamplerCustomNoise(metaclass=IntegratedNode): +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." - @classmethod - def INPUT_TYPES(cls): - 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 { - "required": { - "model": ("MODEL",), - "add_noise": (["enable", "disable"],), - "noise_seed": ( - "INT", - {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}, - ), - "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), - "cfg": ( - "FLOAT", - { - "default": 8.0, - "min": 0.0, - "max": 100.0, - "step": 0.001, - "round": False, - }, - ), - "sampler": ("SAMPLER",), - "scheduler": (restart_normal_schedulers,), - "positive": ("CONDITIONING",), - "negative": ("CONDITIONING",), - "latent_image": ("LATENT",), - "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}), - "end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}), - "return_with_leftover_noise": (["disable", "enable"],), - "segments": ( - "STRING", - { - "default": restart_default_segments, - "multiline": False, - }, - ), - "restart_scheduler": (restart_schedulers,), - "chunked_mode": ("BOOLEAN", {"default": True}), - }, - "optional": { - "custom_noise_opt": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional custom noise input.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - }, - } + INPUT_TYPES = SonarLazyInputTypes(KRestartSamplerCustomNoise_INPUT_TYPES_BUILDER) RETURN_TYPES = ("LATENT", "LATENT") RETURN_NAMES = ("output", "denoised_output") @@ -317,30 +238,20 @@ class KRestartSamplerCustomNoise(metaclass=IntegratedNode): ) -class RestartSamplerCustomNoise(metaclass=IntegratedNode): +class RestartSamplerCustomNoise: DESCRIPTION = "Wrapper used to make another sampler Restart compatible. Allows specifying a custom type for noise added by restarts." - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "sampler": ("SAMPLER",), - "chunked_mode": ("BOOLEAN", {"default": True}), - }, - "optional": { - "custom_noise_opt": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional custom noise input.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - }, - } - 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"): diff --git a/py/nodes/latent_operations.py b/py/nodes/latent_operations.py index 37415bd..ffc7662 100644 --- a/py/nodes/latent_operations.py +++ b/py/nodes/latent_operations.py @@ -1,5 +1,3 @@ -# ruff: noqa: TID252 - from __future__ import annotations import functools @@ -14,7 +12,7 @@ from ..latent_ops import ( SonarLatentOperationNoise, SonarLatentOperationSetSeed, ) -from .base import NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE +from .base import SonarInputTypes, SonarLazyInputTypes from .noise_filters import SonarQuantileFilteredNoiseNode if TYPE_CHECKING: @@ -28,157 +26,96 @@ class SonarApplyLatentOperationCFG(metaclass=IntegratedNode): FUNCTION = "go" - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model": ("MODEL",), - "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).", - }, - ), - "pred_flip_mode": ( - "BOOLEAN", - { - "default": False, - "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.", - }, - ), - "require_uncond": ( - "BOOLEAN", - { - "default": False, - "tooltip": "When enabled, the operation will be skipped if uncond is unavailable. This will also happen if you choose a mode that requires uncond.", - }, - ), - "start_sigma": ( - "FLOAT", - { - "default": -1.0, - "min": -1.0, - "max": 9999.0, - "tooltip": "Sigma when the effect becomes active. You can set a negative value here to use whatever the model's maximum sigma is.", - }, - ), - "end_sigma": ( - "FLOAT", - { - "default": 0.0, - "min": 0.0, - "max": 9999.0, - }, - ), - "blend_mode": ( - tuple(utils.BLENDING_MODES.keys()), - { - "default": "lerp", - "tooltip": "Controls how the output of the latent operation is blended with the original result.", - }, - ), - "blend_strength": ( - "FLOAT", - { - "default": 0.5, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "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, in other words operation_2 sees a full unblended result from operation_1.", - }, - ), - "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.", - }, - ), - "blend_scale_offset": ( - "FLOAT", - { - "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.", - }, - ), - "blend_scale_min": ( - "FLOAT", - { - "default": 0.0, - "min": 0.0, - "max": 1.0, - "tooltip": "Only applies when blend_scale_mode is not none. Minimum value for the blend scale percentage.", - }, - ), - "blend_scale_max": ( - "FLOAT", - { - "default": 1.0, - "min": 0.0, - "max": 1.0, - "tooltip": "Only applies when blend_scale_mode is not none. Maximum value for the blend scale percentage.", - }, - ), - "immediate_blend": ( - "BOOLEAN", - { - "default": False, - }, - ), - }, - "optional": { - "operation_1": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - "operation_2": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - "operation_3": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - "operation_4": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - "operation_5": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - }, - } + 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( @@ -274,7 +211,7 @@ class SonarApplyLatentOperationCFG(metaclass=IntegratedNode): blend_scale_mode = "none" orig_mode = mode - def patch(args: dict) -> torch.Tensor: # noqa: PLR0914 + def patch(args: dict) -> torch.Tensor: nonlocal mode x = args["input"] @@ -402,116 +339,83 @@ class SonarLatentOperationQuantileFilter(SonarQuantileFilteredNoiseNode): norm_factor: float, strategy: str, ): - return ( - SonarLatentOperation( - op=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, - ), - ), + 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, ) + return (SonarLatentOperation(op=lambda latent: qnorm_filter(latent)),) # noqa: PLW0108 + class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode): - DESCRIPTION = "Allows scheduling and other advanced features for latent operations." + 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" - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "operation": ( - "LATENT_OPERATION", - { - "tooltip": "Latent operation to apply.", - }, - ), - "start_sigma": ( - "FLOAT", - { - "default": -1.0, - "min": -1.0, - "max": 9999.0, - "tooltip": "Sigma when the effect becomes active. You can use -1.0 here for no limit.", - }, - ), - "end_sigma": ( - "FLOAT", - { - "default": 0.0, - "min": 0.0, - "max": 9999.0, - }, - ), - "input_multiplier": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "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.", - }, - ), - "output_multiplier": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Flat multiplier on the output from the latent operation. Occurs before blending or calculating the difference.", - }, - ), - "difference_multiplier": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Flat multiplier on the difference or change from the original that the operation performed. Occurs after output_multiplier and before blending applies.", - }, - ), - "blend_mode": ( - tuple(utils.BLENDING_MODES.keys()), - { - "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.", - }, - ), - "blend_strength": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Strength of the blend.", - }, - ), - }, - "optional": { - "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).", - }, - ), - }, - } + INPUT_TYPES = SonarLazyInputTypes( + lambda: SonarInputTypes() + .req_field_operation( + "LATENT_OPERATION", + tooltip="Latent operation to apply.", + ) + .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.", + ) + .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( @@ -526,9 +430,16 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode): blend_mode: str, blend_strength: float, operation_alt=None, + operation_2=None, + operation_3=None, + operation_4=None, + operation_5=None, ) -> tuple[SonarLatentOperationAdvanced]: - if not isinstance(operation, SonarLatentOperation): - operation = SonarLatentOperation(op=operation) + 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, @@ -536,7 +447,7 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode): operation_alt = SonarLatentOperation(op=operation_alt) return ( SonarLatentOperationAdvanced( - op=operation, + ops=operations, op_alt=operation_alt, start_sigma=start_sigma, end_sigma=end_sigma, @@ -556,44 +467,22 @@ class SonarLatentOperationNoiseNode(metaclass=IntegratedNode): FUNCTION = "go" - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "custom_noise": ( - WILDCARD_NOISE, - {"tooltip": f"Custom noise. \n{NOISE_INPUT_TYPES_HINT}"}, - ), - "scale_to_sigma": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Scales the noise to the current sigma.", - }, - ), - "cpu_noise": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether noise is generated on the CPU or GPU. GPU is usually faster but may change seeds for different models of GPU.", - }, - ), - "normalize": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether the generated noise is normalized.", - }, - ), - "lazy_noise_sampler": ( - "BOOLEAN", - { - "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.", - }, - ), - }, - } + 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( @@ -623,22 +512,17 @@ class SonarLatentOperationSetSeedNode(metaclass=IntegratedNode): FUNCTION = "go" - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "operation": ("LATENT_OPERATION",), - "seed": ( - "INT", - { - "default": 0, - "min": 0, - "max": 0xFFFFFFFFFFFFFFFF, - "tooltip": "Seed to set. Note that this is called _every time_ before the operation.", - }, - ), - }, - } + 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( @@ -646,8 +530,15 @@ class SonarLatentOperationSetSeedNode(metaclass=IntegratedNode): *, operation, seed: int, + restore_rng_state: bool, ) -> tuple[SonarLatentOperationSetSeed]: - return (SonarLatentOperationSetSeed(op=operation, seed=seed),) + return ( + SonarLatentOperationSetSeed( + op=operation, + seed=seed, + restore_rng_state=restore_rng_state, + ), + ) NODE_CLASS_MAPPINGS = { diff --git a/py/nodes/misc.py b/py/nodes/misc.py index 5013bc1..d05f040 100644 --- a/py/nodes/misc.py +++ b/py/nodes/misc.py @@ -1,5 +1,3 @@ -# ruff: noqa: TID252 - from __future__ import annotations import functools @@ -12,14 +10,17 @@ 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 ( - NOISE_INPUT_TYPES_HINT, - WILDCARD_NOISE, + NoiseChainInputTypes, SonarCustomNoiseNodeBase, + SonarInputTypes, + SonarLazyInputTypes, SonarNormalizeNoiseNodeMixin, ) @@ -32,96 +33,44 @@ class NoisyLatentLikeNode(metaclass=IntegratedNode): FUNCTION = "go" - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "noise_type": ( - tuple(noise.NoiseType.get_names()), - { - "default": "gaussian", - "tooltip": "Sets the type of noise to generate. Has no effect when the custom_noise_opt input is connected.", - }, - ), - "seed": ( - "INT", - { - "default": 0, - "min": 0, - "max": 0xFFFFFFFFFFFFFFFF, - "tooltip": "Seed to use for generated noise.", - }, - ), - "latent": ( - "LATENT", - { - "tooltip": "Latent used as a reference for generating noise.", - }, - ), - "multiplier": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -10000.0, - "max": 10000.0, - "round": False, - "tooltip": "Multiplier for the strength of the generated noise. Performed after mul_by_sigmas_opt.", - }, - ), - "add_to_latent": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Add the generated noise to the reference latent rather than adding it to an empty latent. Generally should be enabled for img2img workflows.", - }, - ), - "repeat_batch": ( - "INT", - { - "default": 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.", - }, - ), - "cpu_noise": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether noise will be generated on GPU or CPU. Only affects noise types that support GPU generation (maybe only Brownian).", - }, - ), - "normalize": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength before scaling. Generally should be left enabled.", - }, - ), - }, - "optional": { - "custom_noise_opt": ( - WILDCARD_NOISE, - { - "tooltip": f"Allows connecting a custom noise chain. When connected, noise_type has no effect.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "mul_by_sigmas_opt": ( - "SIGMAS", - { - "tooltip": "When connected, will scale the generated noise by the first sigma. Must also connect model_opt to enable.", - }, - ), - "model_opt": ( - "MODEL", - { - "tooltip": "Used when mul_by_sigmas_opt is connected, no effect otherwise.", - }, - ), - }, - } + 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( # noqa: PLR0914 + def go( cls, *, noise_type: str, @@ -213,161 +162,86 @@ class SonarNoiseImageNode(metaclass=IntegratedNode): FUNCTION = "go" - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "noise_type": ( - tuple(NoiseType.get_names()), - { - "default": "gaussian", - "tooltip": "Sets the type of noise to generate. Has no effect when the custom_noise_opt input is connected.", - }, - ), - "seed": ( - "INT", - { - "default": 0, - "min": 0, - "max": 0xFFFFFFFFFFFFFFFF, - "tooltip": "Seed to use for generated noise.", - }, - ), - "image": ( - "IMAGE", - { - "tooltip": "Image noise will be added to.", - }, - ), - "noise_min": ( - "FLOAT", - { - "default": 0.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "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.", - }, - ), - "noise_max": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "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.", - }, - ), - "noise_multiplier": ( - "FLOAT", - { - "default": 0.5, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Multiplier for the strength of the generated noise. This is performed after noise_min/max scaling.", - }, - ), - "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.", - }, - ), - "blend_mode": ( - ("simple_add", *utils.BLENDING_MODES.keys()), - { - "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.", - }, - ), - "blend_strength": ( - "FLOAT", - { - "default": 0.5, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Multiplier for the strength of the generated noise.", - }, - ), - "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.", - }, - ), - "greyscale_mode": ( - "BOOLEAN", - { - "default": False, - "tooltip": "When enabled, generated noise will be averaged so the same amount value is added to all specified channels.", - }, - ), - "pure_noise_mode": ( - "BOOLEAN", - { - "default": False, - "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.", - }, - ), - "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.", - }, - ), - "cpu_noise": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether noise will be generated on GPU or CPU.", - }, - ), - "normalize": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength before scaling. Generally should be left enabled.", - }, - ), - }, - "optional": { - "custom_noise_opt": ( - WILDCARD_NOISE, - { - "tooltip": f"Allows connecting a custom noise chain. When connected, noise_type has no effect.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - }, - } + 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( # noqa: PLR0914 + def go( cls, *, noise_type: str, @@ -551,52 +425,25 @@ class SonarToComfyNOISENode(metaclass=IntegratedNode): CATEGORY = "sampling/custom_sampling/noise" FUNCTION = "go" - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise type to convert.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "seed": ( - "INT", - { - "default": 0, - "min": 0, - "max": 0xFFFFFFFFFFFFFFFF, - "tooltip": "Seed to use for generated noise.", - }, - ), - "cpu_noise": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether noise is generated on CPU or GPU.", - }, - ), - "normalize": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether generated noise is normalized to 1.0 strength.", - }, - ), - "multiplier": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "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).", - }, - ), - }, - } + 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=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.", + ) + .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).", + ), + ) @classmethod def go(cls, *, custom_noise, seed, cpu_noise=True, normalize=True, multiplier=1.0): @@ -614,100 +461,47 @@ class SonarToComfyNOISENode(metaclass=IntegratedNode): 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." - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "sampler": ("SAMPLER",), - "eta": ( - "FLOAT", - { - "default": 1.0, - "step": 0.01, - "max": 1000.0, - "round": False, - "tooltip": "Basically controls the ancestralness of the sampler. When set to 0, you will get a non-ancestral (or SDE) sampler.", - }, - ), - "s_noise": ( - "FLOAT", - { - "default": 1.0, - "step": 0.01, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Multiplier for noise added during ancestral or SDE sampling.", - }, - ), - "s_churn": ( - "FLOAT", - { - "default": 0.0, - "step": 0.01, - "min": -1000.0, - "max": 1000.0, - "round": False, - "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.", - }, - ), - "r": ( - "FLOAT", - { - "default": 0.5, - "step": 0.01, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Used by dpmpp_sde.", - }, - ), - "sde_solver": ( - ("midpoint", "heun"), - { - "tooltip": "Solver used by dpmpp_2m_sde.", - }, - ), - "cpu_noise": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether noise is generated on CPU or GPU.", - }, - ), - "normalize": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether generated noise is normalized to 1.0 strength.", - }, - ), - }, - "optional": { - "noise_type": ( - ("DEFAULT", *NoiseType.get_names()), - { - "default": "DEFAULT", - "tooltip": "Noise type used during ancestral or SDE sampling. Leave blank to use the default for the attached sampler. Only used when the custom noise input is not connected.", - }, - ), - "custom_noise_opt": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional input for custom noise used during ancestral or SDE sampling. When connected, the built-in noise_type selector is ignored.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "yaml_parameters": ( - "STRING", - { - "tooltip": "Allows specifying custom parameters via YAML. Note: When specifying paramaters this way, there is no error checking.", - "placeholder": "# YAML or JSON here", - "dynamicPrompts": False, - "multiline": True, - }, - ), - }, - } + 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" @@ -834,24 +628,13 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode): 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." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - result["optional"] |= { - "custom_noise": ( - WILDCARD_NOISE, - {"tooltip": f"Custom noise. \n{NOISE_INPUT_TYPES_HINT}"}, - ), - } - return result + 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): @@ -878,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/momentum_samplers.py b/py/nodes/momentum_samplers.py index 0605f1f..52b82a7 100644 --- a/py/nodes/momentum_samplers.py +++ b/py/nodes/momentum_samplers.py @@ -1,5 +1,3 @@ -# ruff: noqa: TID252 - from __future__ import annotations from comfy import samplers @@ -15,55 +13,37 @@ from ..sonar import ( SonarEuler, SonarEulerAncestral, ) -from .base import NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE +from .base import SonarInputTypes, SonarLazyInputTypes -class GuidanceConfigNode: +class GuidanceConfigNode(metaclass=IntegratedNode): DESCRIPTION = "Allows specifying extended guidance parameters for Sonar samplers." - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "factor": ( - "FLOAT", - { - "default": 0.01, - "min": -2.0, - "max": 2.0, - "step": 0.001, - "round": False, - "tooltip": "Controls the strength of the guidance. You'll generally want to use fairly low values here.", - }, - ), - "guidance_type": ( - tuple(t.name.lower() for t in GuidanceType), - { - "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.", - }, - ), - "start_step": ( - "INT", - { - "default": 0, - "min": 0, - "tooltip": "First zero-based step the guidance is active.", - }, - ), - "end_step": ( - "INT", - { - "default": 9999, - "min": 0, - "tooltip": "Last zero-based step the guidance is active.", - }, - ), - "latent": ( - "LATENT", - {"tooltip": "Latent to use as a reference for guidance."}, - ), - }, - } + 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" @@ -90,68 +70,44 @@ class GuidanceConfigNode: ) -class SamplerNodeSonarBase(metaclass=IntegratedNode): +class SamplerNodeSonarBase: DESCRIPTION = "Sonar - momentum based sampler node." - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "momentum": ( - "FLOAT", - { - "default": 0.95, - "min": -0.5, - "max": 2.5, - "step": 0.01, - "round": False, - "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.", - }, - ), - "momentum_hist": ( - "FLOAT", - { - "default": 0.75, - "min": -1.5, - "max": 1.5, - "step": 0.01, - "round": False, - "tooltip": "How much of the existing history to leave at each update. 0.75 means keep 75%, mix in 25% of the new result.", - }, - ), - "momentum_init": ( - tuple(t.name for t in HistoryType), - { - "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.", - }, - ), - "direction": ( - "FLOAT", - { - "default": 1.0, - "min": -30.0, - "max": 15.0, - "step": 0.01, - "round": False, - "tooltip": "Multiplier applied to the result of normal sampling.", - }, - ), - "rand_init_noise_type": ( - tuple(NoiseType.get_names(skip=(NoiseType.BROWNIAN,))), - { - "tooltip": "Noise type to use when momentum_init is set to RANDOM.", - }, - ), - }, - "optional": { - "guidance_cfg_opt": ( - "SONAR_GUIDANCE_CFG", - { - "tooltip": "Optional input for extended guidance parameters.", - }, - ), - }, - } + 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_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" @@ -186,52 +142,23 @@ class SamplerNodeSonarEuler(SamplerNodeSonarBase): class SamplerNodeSonarEulerAncestral(SamplerNodeSonarEuler): - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"].update( - { - "s_noise": ( - "FLOAT", - { - "default": 1.0, - "min": -1000.0, - "max": 1000.0, - "step": 0.01, - "round": False, - "tooltip": "Multiplier for noise added during ancestral or SDE sampling.", - }, - ), - "eta": ( - "FLOAT", - { - "default": 1.0, - "min": -1000.0, - "max": 1000.0, - "step": 0.01, - "round": False, - "tooltip": "Basically controls the ancestralness of the sampler. When set to 0, you will get a non-ancestral (or SDE) sampler.", - }, - ), - "noise_type": ( - tuple(NoiseType.get_names()), - { - "tooltip": "Noise type used during ancestral or SDE sampling. Only used when the custom noise input is not connected.", - }, - ), - }, + INPUT_TYPES = SonarLazyInputTypes( + lambda: SonarInputTypes(parent=SamplerNodeSonarEuler) + .req_float_s_noise( + default=1.0, + tooltip="Multiplier for noise added during ancestral or SDE sampling.", ) - result["optional"].update( - { - "custom_noise_opt": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional input for custom noise used during ancestral or SDE sampling. When connected, the built-in noise_type selector is ignored.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - }, + .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.", ) - return result + .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( @@ -270,53 +197,12 @@ class SamplerNodeSonarEulerAncestral(SamplerNodeSonarEuler): ) -class SamplerNodeSonarDPMPPSDE(SamplerNodeSonarEuler): - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"].update( - { - "s_noise": ( - "FLOAT", - { - "default": 1.0, - "min": -1000.0, - "max": 1000.0, - "step": 0.01, - "round": False, - "tooltip": "Multiplier for noise added during ancestral or SDE sampling.", - }, - ), - "eta": ( - "FLOAT", - { - "default": 1.0, - "min": -1000.0, - "max": 1000.0, - "step": 0.01, - "round": False, - "tooltip": "Basically controls the ancestralness of the sampler. When set to 0, you will get a non-ancestral (or SDE) sampler.", - }, - ), - "noise_type": ( - tuple(NoiseType.get_names(default=NoiseType.BROWNIAN)), - { - "tooltip": "Noise type used during ancestral or SDE sampling. Only used when the custom noise input is not connected.", - }, - ), - }, - ) - result["optional"].update( - { - "custom_noise_opt": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional input for custom noise used during ancestral or SDE sampling. When connected, the built-in noise_type selector is ignored.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - }, - ) - return result +class SamplerNodeSonarDPMPPSDE(SamplerNodeSonarEulerAncestral): + INPUT_TYPES = SonarLazyInputTypes( + lambda: SonarInputTypes( + parent=SamplerNodeSonarEulerAncestral, + ).req_selectnoise_noise_type(default="brownian"), + ) @classmethod def get_sampler( diff --git a/py/nodes/noise_filters.py b/py/nodes/noise_filters.py index ca9890a..91ab3a9 100644 --- a/py/nodes/noise_filters.py +++ b/py/nodes/noise_filters.py @@ -1,14 +1,16 @@ -# ruff: noqa: TID252 - from __future__ import annotations +import torch +from comfy import model_management + from .. import noise, utils from ..latent_ops import SonarLatentOperation from ..sonar import SonarGuidanceMixin from .base import ( - NOISE_INPUT_TYPES_HINT, - WILDCARD_NOISE, + NoiseChainInputTypes, + NoiseNoChainInputTypes, SonarCustomNoiseNodeBase, + SonarLazyInputTypes, SonarNormalizeNoiseNodeMixin, ) @@ -16,69 +18,42 @@ from .base import ( class SonarModulatedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): DESCRIPTION = "Custom noise type that allows modulating the output of another custom noise generator." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - result["required"] |= { - "sonar_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Input custom noise to modulate.\n{NOISE_INPUT_TYPES_HINT}", - }, + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_customnoise_sonar_custom_noise(tooltip="Custom noise type to modulate.") + .req_field_modulation_type( + ( + "intensity", + "frequency", + "spectral_signum", + "none", ), - "modulation_type": ( - ( - "intensity", - "frequency", - "spectral_signum", - "none", - ), - { - "tooltip": "Type of modulation to use.", - }, - ), - "dims": ( - "INT", - { - "default": 3, - "min": 1, - "max": 3, - "tooltip": "Dimensions to modulate over. 1 - channels only, 2 - height and width, 3 - both", - }, - ), - "strength": ( - "FLOAT", - { - "default": 2.0, - "min": -100.0, - "max": 100.0, - "step": 0.001, - "round": False, - "tooltip": "Controls the strength of the modulation effect.", - }, - ), - "normalize_result": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the final result is normalized to 1.0 strength.", - }, - ), - "normalize_noise": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - "normalize_ref": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether the reference latent (when present) is normalized to 1.0 strength.", - }, - ), - } - result["optional"] |= {"ref_latent_opt": ("LATENT",)} - return result + tooltip="Type of modulation to use.", + ) + .req_int_dims( + default=3, + min=1, + max=3, + tooltip="Dimensions to modulate over. 1 - channels only, 2 - height and width, 3 - both", + ) + .req_float_strength( + default=2.0, + min=-100.0, + max=100.0, + tooltip="Controls the strength of the modulation effect.", + ) + .req_normalizetristate_normalize_result( + tooltip="Controls whether the final result is normalized to 1.0 strength.", + ) + .req_normalizetristate_normalize_noise( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ) + .req_bool_normalize_ref( + default=True, + tooltip="Controls whether the reference latent (when present) is normalized to 1.0 strength.", + ) + .opt_latent_ref_latent_opt(), + ) @classmethod def get_item_class(cls): @@ -115,48 +90,30 @@ class SonarModulatedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeM class SonarRepeatedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): DESCRIPTION = "Custom noise type that allows caching the output of other custom noise generators." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - result["required"] |= { - "sonar_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input for items to repeat. Note: Unlike most other custom noise nodes, this is treated like a list.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "repeat_length": ( - "INT", - { - "default": 8, - "min": 1, - "max": 100, - "tooltip": "Number of items to cache.", - }, - ), - "max_recycle": ( - "INT", - { - "default": 1000, - "min": 1, - "max": 1000, - "tooltip": "Number of times an individual item will be used before it is replaced with fresh noise.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - "permute": ( - ("enabled", "disabled", "always"), - { - "tooltip": "When enabled, recycled noise will be permuted by randomly flipping it, rolling the channels, etc. If set to always, the noise will be permuted the first time it's used as well.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_customnoise_sonar_custom_noise(tooltip="Custom noise type to modulate.") + .req_int_repeat_length( + default=8, + min=1, + max=100, + tooltip="Number of items to cache.", + ) + .req_int_max_recycle( + default=1000, + min=1, + max=1000, + tooltip="Number of times an individual item will be used before it is replaced with fresh noise.", + ) + .req_normalizetristate_normalize( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ) + .req_field_permute( + ("enabled", "disabled", "always"), + default="enabled", + tooltip="When enabled, recycled noise will be permuted by randomly flipping it, rolling the channels, etc. If set to always, the noise will be permuted the first time it's used as well.", + ), + ) @classmethod def get_item_class(cls): @@ -185,60 +142,33 @@ class SonarRepeatedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMi class SonarScheduledNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): DESCRIPTION = "Custom noise type that allows scheduling the output of other custom noise generators. NOTE: If you don't connect the fallback custom noise input, no noise will be generated outside of the start_percent, end_percent range. I recommend connecting a 1.0 strength Gaussian custom noise node as the fallback." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - result["required"] |= { - "model": ( - "MODEL", - { - "tooltip": "The model input is required to calculate sampling percentages.", - }, - ), - "sonar_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise to use when start_percent and end_percent matches.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "start_percent": ( - "FLOAT", - { - "default": 0.0, - "min": 0.0, - "max": 1.0, - "step": 0.001, - "round": False, - "tooltip": "Time the custom noise becomes active. Note: Sampling percentage where 1.0 indicates 100%, not based on steps.", - }, - ), - "end_percent": ( - "FLOAT", - { - "default": 1.0, - "min": 0.0, - "max": 1.0, - "step": 0.001, - "round": False, - "tooltip": "Time the custom noise effect ends - inclusive, so only sampling percentages greater than this will be excluded. Note: Sampling percentage where 1.0 indicates 100%, not based on steps.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - result["optional"] |= { - "fallback_sonar_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional input for noise to use when outside of the start_percent, end_percent range. NOTE: When not connected, defaults to NO NOISE which is probably not what you want.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_model( + tooltip="The model input is required to calculate sampling percentages.", + ) + .req_customnoise_sonar_custom_noise( + tooltip="Custom noise to use when start_percent and end_percent matches.", + ) + .req_float_start_percent( + default=0.0, + min=0.0, + max=1.0, + tooltip="Time the custom noise becomes active. Note: Sampling percentage where 1.0 indicates 100%, not based on steps.", + ) + .req_float_end_percent( + default=1.0, + min=0.0, + max=1.0, + tooltip="Time the custom noise effect ends - inclusive, so only sampling percentages greater than this will be excluded. Note: Sampling percentage where 1.0 indicates 100%, not based on steps.", + ) + .req_normalizetristate_normalize( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ) + .opt_customnoise_fallback_sonar_custom_noise( + tooltip="Optional input for noise to use when outside of the start_percent, end_percent range. NOTE: When not connected, defaults to NO NOISE which is probably not what you want.", + ), + ) @classmethod def get_item_class(cls): @@ -271,48 +201,28 @@ class SonarScheduledNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeM class SonarCompositeNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): DESCRIPTION = "Custom noise type that allows compositing two other custom noise generators based on a mask." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - result["required"] |= { - "sonar_custom_noise_dst": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input for noise where the mask is not set.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "sonar_custom_noise_src": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input for noise where the mask is set.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "normalize_dst": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether noise generated for dst is normalized to 1.0 strength.", - }, - ), - "normalize_src": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether noise generated for src is normalized to 1.0 strength.", - }, - ), - "normalize_result": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the final result after composition is normalized to 1.0 strength.", - }, - ), - "mask": ( - "MASK", - { - "tooltip": "Mask to use when compositing noise. Where the mask is 1.0, you will get 100% src, where it is 0.75 you will get 75% src and 25% dst. The mask will be rescaled to match the latent size if necessary.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_customnoise_sonar_custom_noise_dst( + tooltip="Custom noise input for noise where the mask is not set.", + ) + .req_customnoise_sonar_custom_noise_src( + tooltip="Custom noise input for noise where the mask is set.", + ) + .req_normalizetristate_normalize_dst( + tooltip="Controls whether noise generated for dst is normalized to 1.0 strength.", + ) + .req_normalizetristate_normalize_src( + tooltip="Controls whether noise generated for src is normalized to 1.0 strength.", + ) + .req_normalizetristate_normalize_result( + tooltip="Controls whether the final result after composition is normalized to 1.0 strength.", + ) + .req_field_mask( + "MASK", + tooltip="Mask to use when compositing noise. Where the mask is 1.0, you will get 100% src, where it is 0.75 you will get 75% src and 25% dst. The mask will be rescaled to match the latent size if necessary.", + ), + ) @classmethod def get_item_class(cls): @@ -343,62 +253,36 @@ class SonarCompositeNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeM class SonarGuidedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): DESCRIPTION = "Custom noise type that mixes a references with another custom noise generator to guide the generation." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - result["required"] |= { - "latent": ( - "LATENT", - { - "tooltip": "Latent to use for guidance.", - }, - ), - "method": ( - ("euler", "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.", - }, - ), - "guidance_factor": ( - "FLOAT", - { - "default": 0.0125, - "min": -100.0, - "max": 100.0, - "step": 0.001, - "round": False, - "tooltip": "Strength of the guidance to apply. Generally should be a relatively slow value to avoid overpowering the generation.", - }, - ), - "normalize_noise": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - "normalize_result": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the final result is normalized to 1.0 strength.", - }, - ), - "normalize_ref": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether the reference latent (when present) is normalized to 1.0 strength.", - }, - ), - } - result["optional"] = { - "sonar_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional custom noise input to combine with the guidance. If you don't attach something here your reference will be combined with zeros.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_latent( + tooltip="Latent to use for guidance.", + ) + .req_field_method( + ("euler", "linear"), + default="euler", + 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_float_guidance_factor( + default=0.0125, + min=-100.0, + max=100.0, + tooltip="Strength of the guidance to apply. Generally should be a relatively slow value to avoid overpowering the generation.", + ) + .req_normalizetristate_normalize_noise( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ) + .req_normalizetristate_normalize_result( + tooltip="Controls whether the final result is normalized to 1.0 strength.", + ) + .req_bool_normalize_ref( + default=True, + tooltip="Controls whether the reference latent (when present) is normalized to 1.0 strength.", + ) + .opt_customnoise_sonar_custom_noise( + tooltip="Optional custom noise input to combine with the guidance. If you don't attach something here your reference will be combined with zeros.", + ), + ) @classmethod def get_item_class(cls): @@ -435,34 +319,21 @@ class SonarGuidedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixi class SonarRandomNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): DESCRIPTION = "Custom noise type that randomly selects between other custom noise items connected to it." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - result["required"] |= { - "sonar_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input for noise items to randomize. Note: Unlike most other custom noise nodes, this is treated like a list.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "mix_count": ( - "INT", - { - "default": 1, - "min": 1, - "max": 100, - "tooltip": "Number of items to select each time noise is generated.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_customnoise_sonar_custom_noise( + tooltip="Custom noise input for noise items to randomize. Note: Unlike most other custom noise nodes, this is treated like a list.", + ) + .req_int_mix_count( + default=1, + min=1, + max=100, + tooltip="Number of items to select each time noise is generated.", + ) + .req_normalizetristate_normalize( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ), + ) @classmethod def get_item_class(cls): @@ -486,32 +357,26 @@ class SonarRandomNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixi class SonarChannelNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): DESCRIPTION = "Custom noise type that uses a different noise generator for each channel. Note: The connected noise items are treated as a list. If you want to blend noise types, you can use something like a SonarBlendedNoise node." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - result["required"] |= { - "sonar_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input for noise items corresponding to each channel. SD1/2x and SDXL use 4 channels, Flux and SD3 use 16. Note: Unlike most other custom noise nodes, this is treated like a list where the noise item furthest from the node corresponds to channel 0.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "insufficient_channels_mode": ( - ("wrap", "repeat", "zero"), - { - "default": "wrap", - "tooltip": "Controls behavior for when there are less noise items connected than channels in the latent. wrap - wraps back to the first noise item, repeat - repeats the last item, zero - fills the channel with zeros (generally not recommended).", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_customnoise_sonar_custom_noise( + tooltip="Custom noise input for noise items corresponding to each channel. SD1/2x and SDXL use 4 channels, Flux and SD3 use 16. Note: Unlike most other custom noise nodes, this is treated like a list where the noise item furthest from the node corresponds to channel 0.", + ) + .req_field_insufficient_channels_mode( + ("wrap", "repeat", "zero"), + default="wrap", + tooltip="Controls behavior for when there are less noise items connected than channels in the latent. wrap - wraps back to the first noise item, repeat - repeats the last item, zero - fills the channel with zeros (generally not recommended).", + ) + .req_int_mix_count( + default=1, + min=1, + max=100, + tooltip="Number of items to select each time noise is generated.", + ) + .req_normalizetristate_normalize( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ), + ) @classmethod def get_item_class(cls): @@ -536,48 +401,28 @@ class SonarChannelNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix class SonarBlendedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): DESCRIPTION = "Custom noise type that allows blending two other noise items." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "noise_2_percent": ( - "FLOAT", - { - "default": 0.5, - "step": 0.001, - "round": False, - "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.", - }, - ), - "blend_mode": ( - tuple(utils.BLENDING_MODES.keys()), - { - "default": "lerp", - "tooltip": "Mode used for blending the two noise types. More modes will be available if ComfyUI-bleh is installed.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength. For weird blend modes, you may want to set this to forced.", - }, - ), - } - result["optional"] |= { - "custom_noise_1": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise. Optional if noise_2 percent is 1.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "custom_noise_2": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise. Optional if noise_2_percent is 0.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + 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 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.", + ) + .req_normalizetristate_normalize( + 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 and custom_noise_mask is not connected..", + ) + .opt_customnoise_custom_noise_2( + 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.", + ), + ) @classmethod def get_item_class(cls): @@ -593,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) @@ -606,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, ) @@ -613,109 +460,74 @@ class SonarBlendedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix class SonarResizedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): DESCRIPTION = "Custom noise type that allows resizing another noise item." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_rescale=False, include_chain=False) - result["required"] |= { - "width": ( - "INT", - { - "default": 1152, - "min": 16, - "max": 1024 * 1024 * 1024, - "step": 8, - "tooltip": "Note: This should almost always be set to a higher value than the image you're actually sampling.", - }, + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_int_width( + default=1152, + min=16, + max=1024 * 1024 * 1024, + step=8, + tooltip="Note: This should almost always be set to a higher value than the image you're actually sampling.", + ) + .req_int_height( + default=1152, + min=16, + max=1024 * 1024 * 1024, + step=8, + tooltip="Note: This should almost always be set to a higher value than the image you're actually sampling.", + ) + .req_field_downscale_strategy( + ("crop", "scale"), + default="crop", + tooltip="Scaling noise is something you'd pretty much only use to create weird effects. For normal workflows, leave this on crop.", + ) + .req_field_initial_reference( + ("prefer_crop", "prefer_scale"), + default="prefer_crop", + tooltip="The initial latent the noise sampler uses as a reference may not match the requested width/height. This setting controls whether to crop or scale. Note: Cropping can only occur when the initial reference is larger than width/height in both dimensions which is unlikely (and not recommended).", + ) + .req_field_crop_mode( + ( + "center", + "top_left", + "top_center", + "top_right", + "center_left", + "center_right", + "bottom_left", + "bottom_center", + "bottom_right", ), - "height": ( - "INT", - { - "default": 1152, - "min": 16, - "max": 1024 * 1024 * 1024, - "step": 8, - "tooltip": "Note: This should almost always be set to a higher value than the image you're actually sampling.", - }, - ), - "downscale_strategy": ( - ("crop", "scale"), - { - "default": "crop", - "tooltip": "Scaling noise is something you'd pretty much only use to create weird effects. For normal workflows, leave this on crop.", - }, - ), - "initial_reference": ( - ("prefer_crop", "prefer_scale"), - { - "default": "prefer_crop", - "tooltip": "The initial latent the noise sampler uses as a reference may not match the requested width/height. This setting controls whether to crop or scale. Note: Cropping can only occur when the initial reference is larger than width/height in both dimensions which is unlikely (and not recommended).", - }, - ), - "crop_mode": ( - ( - "center", - "top_left", - "top_center", - "top_right", - "center_left", - "center_right", - "bottom_left", - "bottom_center", - "bottom_right", - ), - { - "default": "center", - "tooltip": "Note: Crops will have a bias toward the lower number when the size isn't divisible by two. For example, a center crop of size 3 from (0, 1, 2, 3, 4, 5) will result in (1, 2, 3).", - }, - ), - "crop_offset_horizontal": ( - "INT", - { - "default": 0, - "step": 8, - "min": -8000, - "max": 8000, - "tooltip": "This offsets the cropped view by the specified size. Positive values will move it toward the right, negative values will move it toward the left. The offsets will be adjusted to to fit in the available space. For example, if you have crop_mode set to top_right then setting a positive offset isn't going to do anything: it's already as far right as it can go.", - }, - ), - "crop_offset_vertical": ( - "INT", - { - "default": 0, - "step": 8, - "min": -8000, - "max": 8000, - "tooltip": "This offsets the cropped view by the specified size. Positive values will move it toward the bottom, negative values will move it toward the top. The offsets will be adjusted to to fit in the available space. For example, if you have crop_mode set to bottom_right then setting a positive offset isn't going to do anything: it's already as far down as it can go.", - }, - ), - "upscale_mode": ( - utils.UPSCALE_METHODS, - { - "tooltip": "Allows setting the scaling mode when width/height is smaller than the requested size.", - "default": "nearest-exact", - }, - ), - "downscale_mode": ( - utils.UPSCALE_METHODS, - { - "tooltip": "Allows setting the scaling mode when width/height is larger than the requested size and downscale_strategy is set to 'scale'.", - "default": "nearest-exact", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - } - return result + default="center", + tooltip="Note: Crops will have a bias toward the lower number when the size isn't divisible by two. For example, a center crop of size 3 from (0, 1, 2, 3, 4, 5) will result in (1, 2, 3).", + ) + .req_int_crop_offset_horizontal( + default=0, + step=8, + min=-8000, + max=8000, + tooltip="This offsets the cropped view by the specified size. Positive values will move it toward the right, negative values will move it toward the left. The offsets will be adjusted to to fit in the available space. For example, if you have crop_mode set to top_right then setting a positive offset isn't going to do anything: it's already as far right as it can go.", + ) + .req_int_crop_offset_vertical( + default=0, + step=8, + min=-8000, + max=8000, + tooltip="This offsets the cropped view by the specified size. Positive values will move it toward the bottom, negative values will move it toward the top. The offsets will be adjusted to to fit in the available space. For example, if you have crop_mode set to bottom_right then setting a positive offset isn't going to do anything: it's already as far down as it can go.", + ) + .req_selectscalemode_upscale_mode( + tooltip="Allows setting the scaling mode when width/height is smaller than the requested size.", + default="nearest-exact", + ) + .req_selectscalemode_downscale_mode( + tooltip="Allows setting the scaling mode when width/height is larger than the requested size and downscale_strategy is set to 'scale'.", + default="nearest-exact", + ) + .req_normalizetristate_normalize( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ) + .req_customnoise_custom_noise(), + ) @classmethod def get_item_class(cls): @@ -736,11 +548,125 @@ class SonarResizedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix downscale_mode, normalize, custom_noise, + ): + return super().go( + factor, + width=float(width), + height=float(height), + spatial_compression=8, + spatial_mode="absolute", + downscale_strategy=downscale_strategy, + initial_reference=initial_reference, + crop_offset_horizontal=crop_offset_horizontal, + crop_offset_vertical=crop_offset_vertical, + crop_mode=crop_mode, + upscale_mode=upscale_mode, + downscale_mode=downscale_mode, + normalize=normalize, + custom_noise=custom_noise, + ) + + +class SonarResizedNoiseAdvNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixin): + DESCRIPTION = "Custom noise type that allows resizing another noise item. Advanced version of the SonarResizedNoise node." + + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_float_width( + default=32.0, + min=0.0, + tooltip="Note: In absolute mode, this should almost always be set to a higher value than the image you're actually sampling.", + ) + .req_float_height( + default=32.0, + min=0.0, + tooltip="Note: In absolute mode, this should almost always be set to a higher value than the image you're actually sampling.", + ) + .req_field_spatial_mode( + ("relative", "percentage", "absolute"), + default="relative", + tooltip="In relative mode, the sizes control padding. In percentage mode, the values will be interpreted as percentages of the origal size where 1.0 would be 100%, 0.5 would be 50% and so on. In absolute mode, this controls the absolute size.", + ) + .req_int_spatial_compression( + min=1, + default=8, + tooltip="Most image models use 8x spatial compression. When spatial mode is absolute, the sizes will be multiplied by this value. It is ignored in percentage mode.", + ) + .req_field_downscale_strategy( + ("crop", "scale"), + default="crop", + tooltip="Scaling noise is something you'd pretty much only use to create weird effects. For normal workflows, leave this on crop.", + ) + .req_field_initial_reference( + ("prefer_crop", "prefer_scale"), + default="prefer_crop", + tooltip="The initial latent the noise sampler uses as a reference may not match the requested width/height. This setting controls whether to crop or scale. Note: Cropping can only occur when the initial reference is larger than width/height in both dimensions which is unlikely (and not recommended).", + ) + .req_field_crop_mode( + ( + "center", + "top_left", + "top_center", + "top_right", + "center_left", + "center_right", + "bottom_left", + "bottom_center", + "bottom_right", + ), + default="center", + tooltip="Note: Crops will have a bias toward the lower number when the size isn't divisible by two. For example, a center crop of size 3 from (0, 1, 2, 3, 4, 5) will result in (1, 2, 3).", + ) + .req_int_crop_offset_horizontal( + default=0, + tooltip="This offsets the cropped view by the specified size. Positive values will move it toward the right, negative values will move it toward the left. The offsets will be adjusted to to fit in the available space. For example, if you have crop_mode set to top_right then setting a positive offset isn't going to do anything: it's already as far right as it can go.", + ) + .req_int_crop_offset_vertical( + default=0, + tooltip="This offsets the cropped view by the specified size. Positive values will move it toward the bottom, negative values will move it toward the top. The offsets will be adjusted to to fit in the available space. For example, if you have crop_mode set to bottom_right then setting a positive offset isn't going to do anything: it's already as far down as it can go.", + ) + .req_selectscalemode_upscale_mode( + tooltip="Allows setting the scaling mode when width/height is smaller than the requested size.", + default="nearest-exact", + ) + .req_selectscalemode_downscale_mode( + tooltip="Allows setting the scaling mode when width/height is larger than the requested size and downscale_strategy is set to 'scale'.", + default="nearest-exact", + ) + .req_normalizetristate_normalize( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ) + .req_customnoise_custom_noise(), + ) + + @classmethod + def get_item_class(cls): + return noise.ResizedNoise + + def go( + self, + *, + factor: float, + width: float, + height: float, + spatial_mode: str, + spatial_compression: int, + downscale_strategy: str, + initial_reference: str, + crop_offset_horizontal: int, + crop_offset_vertical: int, + crop_mode: str, + upscale_mode: str, + downscale_mode: str, + normalize: str, + custom_noise, ): return super().go( factor, width=width, height=height, + spatial_compression=spatial_compression, + spatial_mode=spatial_mode, downscale_strategy=downscale_strategy, initial_reference=initial_reference, crop_offset_horizontal=crop_offset_horizontal, @@ -756,84 +682,56 @@ class SonarResizedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix class SonarQuantileFilteredNoiseNode(SonarCustomNoiseNodeBase): DESCRIPTION = "Custom noise type that allows filtering noise based on the quantile" - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_chain=False, include_rescale=False) - result["required"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise type to filter.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "quantile": ( - "FLOAT", - { - "default": 0.85, - "min": 0.0, - "max": 1.0, - "step": 0.001, - "round": False, - "tooltip": "When enabled, will normalize generated noise to this quantile (i.e. 0.75 means outliers >75% will be clipped). Set to 1.0 or 0.0 to disable quantile normalization. A value like 0.75 or 0.85 should be reasonable, it really depends on the input and how many of the values are extreme.", - }, - ), - "dim": ( - ("global", "0", "1", "2", "3", "4"), - { - "default": "1", - "tooltip": "Controls what dimensions quantile normalization uses. Dimensions start from 0. Image latents have dimensions: batch, channel, row, column. Video latents have dimensions: batch, channel, frame, row, column.", - }, - ), - "flatten": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether the noise is flattened before quantile normalization. You can try disabling it but they may have a very strong row/column influence.", - }, - ), - "norm_factor": ( - "FLOAT", - { - "default": 1.0, - "min": 0.00001, - "max": 10000.0, - "step": 0.001, - "tooltip": "Multiplier on the input noise just before it is clipped to the quantile min/max. Generally should be left at the default.", - }, - ), - "norm_power": ( - "FLOAT", - { - "default": 0.5, - "min": -10000.0, - "max": 10000.0, - "step": 0.001, - "tooltip": "The absolute value of the noise is raised to this power after it is clipped to the quantile min/max. You can use negative values here, but anything below -0.3 will probably produce pretty strange effects. Generally should be left at the default.", - }, - ), - "normalize_noise": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether the noise source is normalized before quantile filtering occurs.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "default": "disabled", - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength after quantile filtering.", - }, - ), - "strategy": ( - tuple(utils.quantile_handlers.keys()), - { - "default": "clamp", - "tooltip": "Determines how to treat outliers. zero and reverse_zero modes are only useful if you're going to do something like add the result to some other noise. zero will return zero for anything outside the quantile range, reverse_zero only _keeps_ the outliers and zeros everything else.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_customnoise_custom_noise( + tooltip="Custom noise type to filter.", + ) + .req_float_quantile( + default=0.85, + min=0.0, + max=1.0, + step=0.001, + round=False, + tooltip="When enabled, will normalize generated noise to this quantile (i.e. 0.75 means outliers >75% will be clipped). Set to 1.0 or 0.0 to disable quantile normalization. A value like 0.75 or 0.85 should be reasonable, it really depends on the input and how many of the values are extreme.", + ) + .req_field_dim( + ("global", "0", "1", "2", "3", "4"), + default="1", + tooltip="Controls what dimensions quantile normalization uses. Dimensions start from 0. Image latents have dimensions: batch, channel, row, column. Video latents have dimensions: batch, channel, frame, row, column.", + ) + .req_bool_flatten( + default=True, + tooltip="Controls whether the noise is flattened before quantile normalization. You can try disabling it but they may have a very strong row/column influence.", + ) + .req_float_norm_factor( + default=1.0, + min=0.00001, + max=10000.0, + step=0.001, + tooltip="Multiplier on the input noise just before it is clipped to the quantile min/max. Generally should be left at the default.", + ) + .req_float_norm_power( + default=0.5, + min=-10000.0, + max=10000.0, + step=0.001, + tooltip="The absolute value of the noise is raised to this power after it is clipped to the quantile min/max. You can use negative values here, but anything below -0.3 will probably produce pretty strange effects. Generally should be left at the default.", + ) + .req_bool_normalize_noise( + default=False, + tooltip="Controls whether the noise source is normalized before quantile filtering occurs.", + ) + .req_normalizetristate_normalize( + default="disabled", + tooltip="Controls whether the generated noise is normalized to 1.0 strength after quantile filtering.", + ) + .req_field_strategy( + tuple(utils.quantile_handlers.keys()), + default="clamp", + tooltip="Determines how to treat outliers. zero and reverse_zero modes are only useful if you're going to do something like add the result to some other noise. zero will return zero for anything outside the quantile range, reverse_zero only _keeps_ the outliers and zeros everything else.", + ), + ) @classmethod def get_item_class(cls): @@ -870,41 +768,23 @@ class SonarQuantileFilteredNoiseNode(SonarCustomNoiseNodeBase): class SonarShuffledNoiseNode(SonarCustomNoiseNodeBase): DESCRIPTION = "Custom noise type that allows shuffling noise along some dimension" - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_chain=False, include_rescale=False) - result["required"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise type to filter.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "dims": ( - "STRING", - { - "default": "-1", - "tooltip": "Comma separated list of dimensions to shuffle. May be negative to count from the end.", - }, - ), - "flatten": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether to flatten starting from the dimension before the shuffle operation. May be slow as this requires flattening and then reshaping the tensor back to the correct shape. Flattening will occur between the lowest and highest dimension in the list, other dimensions will be ignored. If they are the same, then it will just flatten from the lowest dimension.", - }, - ), - "percentage": ( - "FLOAT", - { - "default": 1.0, - "min": 0.0, - "max": 1.0, - "tooltip": "Percentage of elements to shuffle in the specified dimensions.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_customnoise_custom_noise(tooltip="Custom noise type to filter.") + .req_string_dims( + default="-1", + tooltip="Comma separated list of dimensions to shuffle. May be negative to count from the end.", + ) + .req_bool_flatten( + tooltip="Controls whether to flatten starting from the dimension before the shuffle operation. May be slow as this requires flattening and then reshaping the tensor back to the correct shape. Flattening will occur between the lowest and highest dimension in the list, other dimensions will be ignored. If they are the same, then it will just flatten from the lowest dimension.", + ) + .req_float_percentage( + default=1.0, + min=0.0, + max=1.0, + tooltip="Percentage of elements to shuffle in the specified dimensions.", + ), + ) @classmethod def get_item_class(cls): @@ -933,50 +813,27 @@ class SonarShuffledNoiseNode(SonarCustomNoiseNodeBase): class SonarPatternBreakNoiseNode(SonarCustomNoiseNodeBase): DESCRIPTION = "Custom noise type that allows breaking patterns in the noise with configurable strength" - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES(include_chain=False, include_rescale=False) - result["required"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise type to filter.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "detail_level": ( - "FLOAT", - { - "default": 0.0, - "min": -10000.0, - "max": 10000.0, - "tooltip": "Controls the detail level of the noise when break_pattern is non-zero. No effect when strength is 0.", - }, - ), - "blend_mode": ( - tuple(utils.BLENDING_MODES.keys()), - { - "default": "lerp", - "tooltip": "Function to use for blending original noise with pattern broken noise. If you have ComfyUI-bleh then you will have access to many more blend modes.", - }, - ), - "percentage": ( - "FLOAT", - { - "default": 1.0, - "min": -10000.0, - "max": 10000.0, - "tooltip": "Percentage pattern-broken noise to mix with the original noise. Going outside of 0.0 through 1.0 is unlikely to work well with normal blend modes.", - }, - ), - "restore_scale": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether the original min/max values get preserved. Not sure which is better, it is slightly slower to do this though.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_customnoise_custom_noise(tooltip="Custom noise type to filter.") + .req_float_detail_level( + default=0.0, + tooltip="Controls the detail level of the noise when break_pattern is non-zero. No effect when strength is 0.", + ) + .req_selectblend( + tooltip="Function to use for blending original noise with pattern broken noise. If you have ComfyUI-bleh then you will have access to many more blend modes.", + ) + .req_float_percentage( + default=1.0, + min=0.0, + max=1.0, + tooltip="Percentage of elements to shuffle in the specified dimensions.", + ) + .req_bool_restore_scale( + default=True, + tooltip="Controls whether the original min/max values get preserved. Not sure which is better, it is slightly slower to do this though.", + ), + ) @classmethod def get_item_class(cls): @@ -1040,49 +897,23 @@ yl_scale: 1.0 yh_scales: 1.0 """ - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "normalize_noise": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether the noise source is normalized before wavelet filtering occurs.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "default": "default", - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - result["optional"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional: Custom noise input. If unconnected will default to Gaussian noise.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "custom_noise_high": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional: Custom noise input. If unconnected will use the same noise generator as custom_noise. However, if you do connect it this noise will be used for the high-frequency side of the wavelet.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "yaml_parameters": ( - "STRING", - { - "tooltip": "Allows specifying custom parameters via YAML. Note: When specifying paramaters this way, there is no error checking.", - "placeholder": cls._yaml_placeholder, - "dynamicPrompts": False, - "multiline": True, - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda _yaml_placeholder=_yaml_placeholder: NoiseChainInputTypes() + .req_bool_normalize_noise( + default=False, + tooltip="Controls whether the noise source is normalized before wavelet filtering occurs.", + ) + .req_normalizetristate_normalize( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ) + .opt_customnoise_custom_noise( + tooltip="Optional: Custom noise input. If unconnected will default to Gaussian noise.", + ) + .opt_customnoise_custom_noise_high( + tooltip="Optional: Custom noise input. If unconnected will use the same noise generator as custom_noise. However, if you do connect it this noise will be used for the high-frequency side of the wavelet.", + ) + .opt_yaml(placeholder=_yaml_placeholder), + ) @classmethod def get_item_class(cls): @@ -1120,89 +951,61 @@ class SonarScatternetFilteredNoiseNode( ): DESCRIPTION = "Custom noise type that allows filtering noise using a scatternet (basically wavelets). Requires the pytorch_wavelets package to be installed in your Python environment. Can be used to do stuff like take the higher frequency components of a very low-frequency noise type such as Pyramid. Currently only works with 4D latents." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "padding_mode": ( - "STRING", - { - "default": "symmetric", - "tooltip": "This is just passed to the pytorch_wavelets scatternet constructor. Valid padding modes that I know of (second order only supports symmetric and zero): symmetric, reflect, zero, periodization, constant, replicate, periodic", - }, + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_string_padding_mode( + default="symmetric", + tooltip="This is just passed to the pytorch_wavelets scatternet constructor. Valid padding modes that I know of (second order only supports symmetric and zero): symmetric, reflect, zero, periodization, constant, replicate, periodic", + ) + .req_bool_use_symmetric_filter( + default=False, + tooltip="Slower, but possibly higher quality.", + ) + .req_float_magbias( + default=1e-02, + min=-1000.0, + max=1000.0, + tooltip="Magnitude bias. Changing it doesn't seem to affect anything, but you can try.", + ) + .req_float_output_offset( + default=0.0, + min=-100000.0, + max=100000.0, + tooltip="Controls where the output starts. The beginning is the low frequency bands, the end is high frequencies. If less than 1 (positive or negative) it will be treated as a percentage into the dimension. Negative values count from the end.", + ) + .req_field_output_mode( + ( + "channels_adjusted", + "flat_adjusted", + "channels", + "flat", + "channels_scaled", + "flat_scaled", ), - "use_symmetric_filter": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Slower, but possibly higher quality.", - }, - ), - "magbias": ( - "FLOAT", - { - "default": 1e-02, - "min": -1000.0, - "max": 1000.0, - "tooltip": "Magnitude bias. Changing it doesn't seem to affect anything, but you can try.", - }, - ), - "output_offset": ( - "FLOAT", - { - "default": 0.0, - "min": -100000.0, - "max": 100000.0, - "tooltip": "Controls where the output starts. The beginning is the low frequency bands, the end is high frequencies. If less than 1 (positive or negative) it will be treated as a percentage into the dimension. Negative values count from the end.", - }, - ), - "output_mode": ( - ("channels_adjusted", "flat_adjusted", "channels", "flat"), - { - "default": "channels_adjusted", - "tooltip": "The normal scatternet reduces the spatial dimensions 2x, the second order one 4x. The adjusted modes will generate larger noise (in the spatial dimensions) to compensate, this is slower but gives you a lot more room to work with. Modes that start with channels will index along the channel dimension, otherwise the indexing will be flat (after the batch dimension). Note: I recommend channels_adjusted mode, it's very possible the offset indexing math is wrong for other modes.", - }, - ), - "scatternet_order": ( - "INT", - { - "default": 1, - "min": -3, - "max": 3, - "tooltip": "Each order increases the number of channels exponentially. You can use a primitive node to bypass the limit of 3 here if you're a crazy person, the code will handle any value but you're very likely to die of old age or run out of VRAM or both if you go above 3 (and even that is stretching it). You can set this to 0 to disable scatternet filtering quickly. Negative values are the same as positive ones here with one exception: there's a specialized 2nd order scatternet which will be used by default for order 2, however it may not support the normal parameters (like padding modes). Use -2 here if you just want to stack two normal scatternet layers instead.", - }, - ), - "per_channel_scatternet": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Runs the scatternet on each channel separately. May be very slow. Models like SDXL use 4 channels, models like Flux have 16. Enabling this may help with non-adjusted output modes.", - }, - ), - "normalize_noise": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether the noise source is normalized before scatternet filtering occurs.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "default": "default", - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - result["optional"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional: Custom noise input. If unconnected will default to Gaussian noise.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - } - return result + default="channels_adjusted", + tooltip="The normal scatternet reduces the spatial dimensions 2x, the second order one 4x. The adjusted modes will generate larger noise (in the spatial dimensions) to compensate, this is slower but gives you a lot more room to work with. The scaled modes will just scale the noise to compensate (likely doesn't work well). Modes that start with channels will index along the channel dimension, otherwise the indexing will be flat (after the batch dimension). Note: I recommend channels_adjusted mode, it's very possible the offset indexing math is wrong for other modes.", + ) + .req_int_scatternet_order( + default=1, + min=-3, + max=3, + tooltip="Each order increases the number of channels exponentially. You can use a primitive node to bypass the limit of 3 here if you're a crazy person, the code will handle any value but you're very likely to die of old age or run out of VRAM or both if you go above 3 (and even that is stretching it). You can set this to 0 to disable scatternet filtering quickly. Negative values are the same as positive ones here with one exception: there's a specialized 2nd order scatternet which will be used by default for order 2, however it may not support the normal parameters (like padding modes). Use -2 here if you just want to stack two normal scatternet layers instead.", + ) + .req_bool_per_channel_scatternet( + default=False, + tooltip="Runs the scatternet on each channel separately. May be very slow. Models like SDXL use 4 channels, models like Flux have 16. Enabling this may help with non-adjusted output modes.", + ) + .req_bool_normalize_noise( + default=False, + tooltip="Controls whether the noise source is normalized before scatternet filtering occurs.", + ) + .req_normalizetristate_normalize( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ) + .opt_customnoise_custom_noise( + tooltip="Optional: Custom noise input. If unconnected will default to Gaussian noise.", + ), + ) @classmethod def get_item_class(cls): @@ -1248,99 +1051,62 @@ class SonarRippleFilteredNoiseNode( "Custom noise filter that allows applying scaling based on a wave (sin or cos)." ) - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input. \n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "mode": ( - ("sin", "cos", "sin_copysign", "cos_copysign"), - { - "default": "cos", - "tooltip": "Function to use for rippling. The copysign variations are not recommended, they will force the noise to the sign of the wave (whether it's above or below the midline) which has an extremely strong effect. If you want to try it, use something like a 1:16 ratio or higher with normal noise.", - }, - ), - "dim": ( - "INT", - { - "default": -1, - "min": -100, - "max": 100, - "tooltip": "Dimension to use for the ripple effect. Negative dimensions count from the end where -1 is the last dimension.", - }, - ), - "flatten": ( - "BOOLEAN", - { - "default": False, - "tooltip": "When enabled, the noise will be flattened starting from (and including) the specified dimension.", - }, - ), - "offset": ( - "FLOAT", - { - "default": 0.0, - "min": -10000, - "max": 10000.0, - "tooltip": "Simple addition to the base value used for the wave.", - }, - ), - "roll": ( - "FLOAT", - { - "default": 0.0, - "min": -10000, - "max": 10000.0, - "tooltip": "Rolls the wave by this many elements each time the noise generator is called. Negative values roll backward.", - }, - ), - "amplitude_high": ( - "FLOAT", - { - "default": 0.25, - "min": -10000, - "max": 10000.0, - "tooltip": "Scale for noise at the highest point of the wave. This adds to the base value (respecting sign). For example, if set to 0.25 you will get noise * 1.25 at that point. It's also possible to use negative values, -0.25 will result in noise * -1.25.", - }, - ), - "amplitude_low": ( - "FLOAT", - { - "default": 0.15, - "min": -10000, - "max": 10000.0, - "tooltip": "Scale for noise at the lowest point of the wave. This subtracts from the base value (respecting sign). For example, if set to 0.25 you will get noise * 0.75 at that point. It's also possible to use negative values, -0.25 will result in noise * -0.75.", - }, - ), - "period": ( - "FLOAT", - { - "default": 3.0, - "min": -10000, - "max": 10000.0, - "tooltip": "Number of oscillations along the specified dimension.", - }, - ), - "normalize_noise": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether the noise source is normalized before wavelet filtering occurs.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_customnoise_custom_noise() + .req_field_mode( + ("sin", "cos", "sin_copysign", "cos_copysign"), + default="cos", + tooltip="Function to use for rippling. The copysign variations are not recommended, they will force the noise to the sign of the wave (whether it's above or below the midline) which has an extremely strong effect. If you want to try it, use something like a 1:16 ratio or higher with normal noise.", + ) + .req_int_dim( + default=-1, + min=-100, + max=100, + tooltip="Dimension to use for the ripple effect. Negative dimensions count from the end where -1 is the last dimension.", + ) + .req_bool_flatten( + default=False, + tooltip="When enabled, the noise will be flattened starting from (and including) the specified dimension.", + ) + .req_float_offset( + default=0.0, + min=-10000, + max=10000.0, + tooltip="Simple addition to the base value used for the wave.", + ) + .req_float_roll( + default=0.0, + min=-10000, + max=10000.0, + tooltip="Rolls the wave by this many elements each time the noise generator is called. Negative values roll backward.", + ) + .req_float_amplitude_high( + default=0.25, + min=-10000, + max=10000.0, + tooltip="Scale for noise at the highest point of the wave. This adds to the base value (respecting sign). For example, if set to 0.25 you will get noise * 1.25 at that point. It's also possible to use negative values, -0.25 will result in noise * -1.25.", + ) + .req_float_amplitude_low( + default=0.15, + min=-10000, + max=10000.0, + tooltip="Scale for noise at the lowest point of the wave. This subtracts from the base value (respecting sign). For example, if set to 0.25 you will get noise * 0.75 at that point. It's also possible to use negative values, -0.25 will result in noise * -0.75.", + ) + .req_float_period( + default=3.0, + min=-10000, + max=10000.0, + tooltip="Number of oscillations along the specified dimension.", + ) + .req_bool_normalize_noise( + default=False, + tooltip="Controls whether the noise source is normalized before wavelet filtering occurs.", + ) + .req_normalizetristate_normalize( + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ), + ) @classmethod def get_item_class(cls): @@ -1388,82 +1154,71 @@ class SonarNormalizeNoiseToScaleNode( ): DESCRIPTION = "Custom noise type that allows precisely controling noise normalization. The default range of -4.5 to 4.5 is roughly what you'd get from 10,000 items of Gaussian noise." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input. \n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "min_negative_value": ( - "FLOAT", - { - "default": -4.5, - "min": -10000.0, - "max": 10000.0, - "tooltip": "In simple mode, this is just the lowest value in the range (and can be positive, despite the name). In advanced mode, this controls the minimum negative value. If you set it to 0 or higher then normalization will leave negative values alone.", - }, - ), - "max_negative_value": ( - "FLOAT", - { - "default": 0.0, - "min": -10000.0, - "max": 10000.0, - "tooltip": "Not used in simple mode. In advanced mode, this controls the maximum negative value. If you set it to 0 or higher, a maximum negative value will be automatically determined from negative value closest (but not equal to) zero.", - }, - ), - "min_positive_value": ( - "FLOAT", - { - "default": 0.0, - "min": -10000.0, - "max": 10000.0, - "tooltip": "Not used in simple mode. In advanced mode, this controls the minmum positive value. If you set it to 0 or lower, a minimum positive value will be automatically determined from positive value closest (but not equal to) zero.", - }, - ), - "max_positive_value": ( - "FLOAT", - { - "default": 4.5, - "min": -10000.0, - "max": 10000.0, - "tooltip": "In simple mode, this is just the highest value in the range (and can be negative, despite the name). In advanced mode, this controls the maximum positive value. If you set it to 0 or lower then normalization will leave positive values alone.", - }, - ), - "mode": ( - ("simple", "advanced"), - { - "default": "simple", - "tooltip": "There are several modes:\nsimple: The noise will be rebalanced to be in between min_negative_value and max_positive_value. Though it sounds weird, you don't need to respect the positive/negative in the names. It is just treated as a simple range.\nadvanced: Positive and negative values in the noise are separately rebalanced to be between the specified ranges. If you set max_negative_value to something positive or min_positive_value to something negative this will automatically determine whatever the closest value to zero is for each sign. Additionally, if you set max_positive_value to something negative or min_negative_value to something positive then values for that sign will be left alone.", - }, - ), - "dims": ( - "STRING", - { - "default": "-3, -2, -1", - "tooltip": "A comma separated list of dimensions which can be negative to count from the end. This behaves differently in advanced mode: If left blank, normalization will be global. If set to anything, normalization will be over each batch item separately. The actual values of the dimensions are ignored in advanced mode currently.", - }, - ), - "normalize_noise": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether the noise source is normalized before wavelet filtering occurs.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "default": "disabled", - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_customnoise_custom_noise() + .req_float_min_negative_value( + default=-4.5, + min=-10000.0, + max=10000.0, + tooltip="In simple mode, this is just the lowest value in the range (and can be positive, despite the name). In advanced mode, this controls the minimum negative value. If you set it to 0 or higher then normalization will leave negative values alone.", + ) + .req_float_max_negative_value( + default=0.0, + min=-10000.0, + max=10000.0, + tooltip="Not used in simple mode. In advanced mode, this controls the maximum negative value. If you set it to 0 or higher, a maximum negative value will be automatically determined from negative value closest (but not equal to) zero.", + ) + .req_float_min_positive_value( + default=0.0, + min=-10000.0, + max=10000.0, + tooltip="Not used in simple mode. In advanced mode, this controls the minmum positive value. If you set it to 0 or lower, a minimum positive value will be automatically determined from positive value closest (but not equal to) zero.", + ) + .req_float_max_positive_value( + default=4.5, + min=-10000.0, + max=10000.0, + tooltip="In simple mode, this is just the highest value in the range (and can be negative, despite the name). In advanced mode, this controls the maximum positive value. If you set it to 0 or lower then normalization will leave positive values alone.", + ) + .req_field_mode( + ("simple", "advanced"), + default="simple", + tooltip="There are several modes:\nsimple: The noise will be rebalanced to be in between min_negative_value and max_positive_value. Though it sounds weird, you don't need to respect the positive/negative in the names. It is just treated as a simple range.\nadvanced: Positive and negative values in the noise are separately rebalanced to be between the specified ranges. If you set max_negative_value to something positive or min_positive_value to something negative this will automatically determine whatever the closest value to zero is for each sign. Additionally, if you set max_positive_value to something negative or min_negative_value to something positive then values for that sign will be left alone.", + ) + .req_string_dims( + default="-3, -2, -1", + tooltip="A comma separated list of dimensions which can be negative to count from the end. This behaves differently in advanced mode: If left blank, normalization will be global. If set to anything, normalization will be over each batch item separately. The actual values of the dimensions are ignored in advanced mode currently.", + ) + .req_string_std_dims( + default="-3, -2, -1", + tooltip="A comma separated list of dimensions which can be negative to count from the end.", + ) + .req_float_std_multiplier( + default=1.0, + min=-10000.0, + max=10000.0, + tooltip="Multiplier on the distance of the std from 1.0. The noise will be divided by this. You can set it to 1.0 to skip the division. When enabled, the division occurs before the final normalize and scaling and after the min/max value parameters are applied.", + ) + .req_string_mean_dims( + default="-3, -2, -1", + tooltip="A comma separated list of dimensions which can be negative to count from the end.", + ) + .req_float_mean_multiplier( + default=1.0, + min=-10000.0, + max=10000.0, + tooltip="Multiplier on the mean of the noise. The mean will be subtracted from the noise if it's not 0. This occurs before the final normalize and scaling and after the min/max value parameters are applied.", + ) + .req_bool_normalize_noise( + default=False, + tooltip="Controls whether the noise source is normalized immediately after generation.", + ) + .req_normalizetristate_normalize( + default="disabled", + tooltip="Controls whether the generated noise is normalized to 1.0 strength. Enabling this does the same thing as the default mean/std settings.", + ), + ) @classmethod def get_item_class(cls): @@ -1481,6 +1236,10 @@ class SonarNormalizeNoiseToScaleNode( max_positive_value: float, mode: str, dims: str, + std_dims: str, + std_multiplier: float, + mean_dims: str, + mean_multiplier: float, normalize_noise: bool, custom_noise=None, sonar_custom_noise_opt=None, @@ -1495,6 +1254,14 @@ class SonarNormalizeNoiseToScaleNode( max_positive_value=max_positive_value, mode=mode, dims=() if not dims.strip() else tuple(int(i) for i in dims.split(",")), + std_dims=() + if not std_dims.strip() + else tuple(int(i) for i in dims.split(",")), + std_multiplier=std_multiplier, + mean_dims=() + if not mean_dims.strip() + else tuple(int(i) for i in dims.split(",")), + mean_multiplier=mean_multiplier, normalize=self.get_normalize(normalize), normalize_noise=normalize_noise, noise=custom_noise, @@ -1507,65 +1274,34 @@ class SonarPerDimNoiseNode( ): DESCRIPTION = "Custom noise type that allows calling the noise sampler multiple times along a dimension. Can be useful for stuff like moving slices of 3D Perlin noise into the batch dimension." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input. \n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "dim": ( - "INT", - { - "default": 0, - "min": -100, - "max": 100, - "tooltip": "Dimension to use. The default usually corresponds to the batch. Be careful using dimensions above 1 as those tend to be spatial and you might end up calling a slow noise sampler hundreds of times.", - }, - ), - "shrink_dim": ( - "BOOLEAN", - { - "default": False, - "tooltip": "When enabled, the reference latent will be chunk_size in the specified dimension. When disabled, noise will be generated according to the initial latent size and then sliced along the specified dimension. Enabling it should be considerably faster/more memory efficient but may not work well for some noise types.", - }, - ), - "chunk_size": ( - "INT", - { - "default": 1, - "min": 1, - "max": 10000, - "tooltip": "Can be used to control how many times the noise sampler is called. For example, if you have dim=0, chunk_size=2 and are dealing with a batch of 4, this will call the noise sampler twice, taking the first two items from the first call and the last two items from the second call.", - }, - ), - # "offset": ( - # "INT", - # { - # "default": 0, - # "min": -10000, - # "max": 10000, - # }, - # ), - "normalize_noise": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether the noise source is normalized initially.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "default": "disabled", - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_customnoise_custom_noise() + .req_int_dim( + default=0, + min=-100, + max=100, + tooltip="Dimension to use. The default usually corresponds to the batch. Be careful using dimensions above 1 as those tend to be spatial and you might end up calling a slow noise sampler hundreds of times.", + ) + .req_bool_shrink_dim( + default=False, + tooltip="When enabled, the reference latent will be chunk_size in the specified dimension. When disabled, noise will be generated according to the initial latent size and then sliced along the specified dimension. Enabling it should be considerably faster/more memory efficient but may not work well for some noise types.", + ) + .req_int_chunk_size( + default=1, + min=1, + max=10000, + tooltip="Can be used to control how many times the noise sampler is called. For example, if you have dim=0, chunk_size=2 and are dealing with a batch of 4, this will call the noise sampler twice, taking the first two items from the first call and the last two items from the second call.", + ) + .req_bool_normalize_noise( + default=False, + tooltip="Controls whether the noise source is normalized initially.", + ) + .req_normalizetristate_normalize( + default="disabled", + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ), + ) @classmethod def get_item_class(cls): @@ -1591,7 +1327,7 @@ class SonarPerDimNoiseNode( sonar_custom_noise_opt=sonar_custom_noise_opt, dim=dim, shrink_dim=shrink_dim, - # offset=offset, + offset=0, chunk_size=chunk_size, normalize=self.get_normalize(normalize), normalize_noise=normalize_noise, @@ -1605,64 +1341,38 @@ class SonarLatentOperationFilteredNoiseNode( ): DESCRIPTION = "Custom noise type that allows filtering noise with a LATENT_OPERATION. If you connect more than one, the operations will be run in sequence." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Custom noise input. \n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "normalize_noise": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether the noise source is normalized initially.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "default": "disabled", - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.", - }, - ), - } - result["optional"] |= { - "operation_1": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - "operation_2": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - "operation_3": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - "operation_4": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - "operation_5": ( - "LATENT_OPERATION", - { - "tooltip": "Optional LATENT_OPERATION. The operations will be applied in sequence.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_customnoise_custom_noise() + .req_bool_normalize_noise( + default=False, + tooltip="Controls whether the noise source is normalized initially.", + ) + .req_normalizetristate_normalize( + default="disabled", + tooltip="Controls whether the generated noise is normalized to 1.0 strength.", + ) + .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.", + ), + ) @classmethod def get_item_class(cls): @@ -1699,23 +1409,147 @@ class SonarLatentOperationFilteredNoiseNode( ) +class SonarCustomNoiseParametersNode( + SonarCustomNoiseNodeBase, + SonarNormalizeNoiseNodeMixin, +): + DESCRIPTION = "Custom noise type that allows setting parameters like dtype or forking the RNG." + + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseNoChainInputTypes() + .req_customnoise_custom_noise() + .req_int_rng_state_offset( + default=0, + min=0, + tooltip="In other words, seed. Avoiding using the word seed here to suppress ComfyUI's annoying default behavior. If you want stuff like auto-increment you can connect an INT primitive node.", + ) + .req_field_rng_offset_mode( + ("disabled", "override", "add"), + default="disabled", + tooltip="Controls the seed passed to the noise sampler and also seeding when rng_mode is set to separate. Most noise samplers don't care about the seed so this generally will only have an effect in when rng_mode is set to separate.", + ) + .req_field_rng_mode( + ("default", "separate", "fork"), + default="default", + tooltip="default mode doesn't do anything special. separate mode creates a generator and saves/restores the state when generating noise (also includes the Python random module). fork uses the existing RNG state (for both Torch and Python random module) but restores it to whatever it was before the custom noise was called.", + ) + .req_bool_frames_to_channels( + tooltip="Only applicable for 5D latents (video models). Will move the frame dimension into channels, may be necessary if a noise type can't deal with 5D latents directly. It's safe to enable this for all models.", + ) + .req_bool_ensure_square_aspect_ratio( + tooltip="Will rearrange the height/width sizes to be square, padding with zeros if necessary. May help some noise types work better with extreme aspect ratios, can also deal with 3D (1 spatial dimension) latents.", + ) + .req_bool_fix_invalid( + tooltip="Replaces any NaNs or infinite values with 0.", + ) + .req_field_override_dtype( + ( + "default", + "float64", + "float32", + "float16", + "bfloat16", + "float8_e4m3fn", + "float8_e4m3fnuz", + "float8_e5m2", + "float8_e5m2fnuz", + "float8_e8m0fnu", + "int64", + "int32", + "int16", + "int8", + ), + default="default", + tooltip="Can be used to override the dtype the noise is generated with. Not all noise generators support all types. I don't recommend using the int or float8 types. Probably the most useful override is float64.", + ) + .req_field_override_device( + ("default", "cpu", "gpu"), + default="default", + tooltip="default just uses whatever device normally would be used. gpu will use ComfyUI's default GPU device and also toggle the cpu_noise flag off. cpu will use the CPU device and toggle the cpu_noise flag on.", + ) + .req_normalizetristate_normalize(), + ) + + @classmethod + def get_item_class(cls): + return noise.CustomNoiseParametersNoise + + def go( + self, + *, + factor, + rng_state_offset: int, + rng_offset_mode: str, + rng_mode: str, + frames_to_channels: bool, + ensure_square_aspect_ratio: bool, + fix_invalid: bool, + override_dtype: str, + override_device: str, + normalize: str, + custom_noise: object, + ): + valid_dtypes = { + "default", + "float64", + "float32", + "float16", + "bfloat16", + "float8_e4m3fn", + "float8_e4m3fnuz", + "float8_e5m2", + "float8_e5m2fnuz", + "float8_e8m0fnu", + "int64", + "int32", + "int16", + "int8", + } + dt = getattr(torch, override_dtype, None) + if override_dtype not in valid_dtypes or ( + override_dtype != "default" and dt is None + ): + raise ValueError("Bad dtype, may not be supported by your PyTorch version") + if override_device == "default": + device = None + elif override_device == "cpu": + device = "cpu" + elif override_device == "gpu": + device = model_management.get_torch_device() + return super().go( + factor, + rng_state_offset=rng_state_offset, + rng_offset_mode=rng_offset_mode, + rng_mode=rng_mode, + frames_to_channels=frames_to_channels, + ensure_square_aspect_ratio=ensure_square_aspect_ratio, + fix_invalid=fix_invalid, + override_dtype=dt, + override_device=device, + normalize=normalize, + noise=custom_noise, + ) + + NODE_CLASS_MAPPINGS = { - "SonarCompositeNoise": SonarCompositeNoiseNode, - "SonarModulatedNoise": SonarModulatedNoiseNode, - "SonarRepeatedNoise": SonarRepeatedNoiseNode, - "SonarScheduledNoise": SonarScheduledNoiseNode, - "SonarGuidedNoise": SonarGuidedNoiseNode, - "SonarRandomNoise": SonarRandomNoiseNode, - "SonarShuffledNoise": SonarShuffledNoiseNode, - "SonarPatternBreakNoise": SonarPatternBreakNoiseNode, - "SonarChannelNoise": SonarChannelNoiseNode, "SonarBlendedNoise": SonarBlendedNoiseNode, - "SonarResizedNoise": SonarResizedNoiseNode, - "SonarWaveletFilteredNoise": SonarWaveletFilteredNoiseNode, - "SonarRippleFilteredNoise": SonarRippleFilteredNoiseNode, - "SonarQuantileFilteredNoise": SonarQuantileFilteredNoiseNode, - "SonarNormalizeNoiseToScale": SonarNormalizeNoiseToScaleNode, - "SonarPerDimNoise": SonarPerDimNoiseNode, - "SonarScatternetFilteredNoise": SonarScatternetFilteredNoiseNode, + "SonarChannelNoise": SonarChannelNoiseNode, + "SonarCompositeNoise": SonarCompositeNoiseNode, + "SonarCustomNoiseParameters": SonarCustomNoiseParametersNode, + "SonarGuidedNoise": SonarGuidedNoiseNode, "SonarLatentOperationFilteredNoise": SonarLatentOperationFilteredNoiseNode, + "SonarModulatedNoise": SonarModulatedNoiseNode, + "SonarNormalizeNoiseToScale": SonarNormalizeNoiseToScaleNode, + "SonarPatternBreakNoise": SonarPatternBreakNoiseNode, + "SonarPerDimNoise": SonarPerDimNoiseNode, + "SonarQuantileFilteredNoise": SonarQuantileFilteredNoiseNode, + "SonarRandomNoise": SonarRandomNoiseNode, + "SonarRepeatedNoise": SonarRepeatedNoiseNode, + "SonarResizedNoise": SonarResizedNoiseNode, + "SonarResizedNoiseAdv": SonarResizedNoiseAdvNode, + "SonarRippleFilteredNoise": SonarRippleFilteredNoiseNode, + "SonarScatternetFilteredNoise": SonarScatternetFilteredNoiseNode, + "SonarScheduledNoise": SonarScheduledNoiseNode, + "SonarShuffledNoise": SonarShuffledNoiseNode, + "SonarWaveletFilteredNoise": SonarWaveletFilteredNoiseNode, } diff --git a/py/nodes/noise_types.py b/py/nodes/noise_types.py index 99dd9e0..b2f2f2d 100644 --- a/py/nodes/noise_types.py +++ b/py/nodes/noise_types.py @@ -1,15 +1,13 @@ -# ruff: noqa: TID252 - from __future__ import annotations import torch from .. import noise, utils -from ..noise_generation import DistroNoiseGenerator +from ..noise_generation import DistroNoiseGenerator, VoronoiNoiseGenerator from .base import ( - NOISE_INPUT_TYPES_HINT, - WILDCARD_NOISE, + NoiseChainInputTypes, SonarCustomNoiseNodeBase, + SonarLazyInputTypes, SonarNormalizeNoiseNodeMixin, ) @@ -19,50 +17,33 @@ class SonarAdvancedPyramidNoiseNode(SonarCustomNoiseNodeBase): "Custom noise type that allows specifying parameters for Pyramid variants." ) - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "variant": ( - ( - "highres_pyramid", - "pyramid", - "pyramid_old", - ), - { - "tooltip": "Sets the Pyramid noise variant to generate.", - "default": "highres_pyramid", - }, + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_field_variant( + ( + "highres_pyramid", + "pyramid", + "pyramid_old", ), - "iterations": ( - "INT", - { - "default": -1, - "min": -1, - "max": 8, - "tooltip": "When set to -1 will use the variant default.", - }, - ), - "discount": ( - "FLOAT", - { - "default": 0.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "When set to 0 will use the variant default.", - }, - ), - "upscale_mode": ( - ("default", *utils.UPSCALE_METHODS), - { - "tooltip": "Allows setting the scaling mode for Pyramid noise. Leave on default to use the variant default.", - "default": "default", - }, - ), - } - return result + default="highres_pyramid", + tooltip="Sets the Pyramid noise variant to generate.", + ) + .req_int_iterations( + default=-1, + min=-1, + max=8, + tooltip="When set to -1 will use the variant default.", + ) + .req_float_discount( + default=0.0, + tooltip="When set to 0 will use the variant default.", + ) + .req_selectscalemode_upscale_mode( + insert_modes=("default",), + default="default", + tooltip="Allows setting the scaling mode for Pyramid noise. Leave on default to use the variant default.", + ), + ) @classmethod def get_item_class(cls): @@ -93,63 +74,29 @@ class SonarAdvancedPyramidNoiseNode(SonarCustomNoiseNodeBase): class SonarAdvanced1fNoiseNode(SonarCustomNoiseNodeBase): DESCRIPTION = "Custom noise type that allows specifying parameters for 1f (pink, green, etc) variants." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "alpha": ( - "FLOAT", - { - "default": 0.25, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Similar to the advanced power noise node, positive values increase low frequencies (with colorful effects), negative values increase high frequencies.", - }, - ), - "k": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Currently no description of exactly what it does, it's just another knob you can try turning for a different effect.", - }, - ), - "vertical_factor": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Vertical frequency scaling factor.", - }, - ), - "horizontal_factor": ( - "FLOAT", - { - "default": 1.0, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Horizontal frequency scaling factor.", - }, - ), - "use_sqrt": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether to sqrt when dividing the FFT. Negative hfac/wfac won't work when enabled. Turning it off seems to make the parameters have a much stronger effect.", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_float_alpha( + default=0.25, + tooltip="Similar to the advanced power noise node, positive values increase low frequencies (with colorful effects), negative values increase high frequencies.", + ) + .req_float_k( + default=1.0, + tooltip="Currently no description of exactly what it does, it's just another knob you can try turning for a different effect.", + ) + .req_float_vertical_factor( + default=1.0, + tooltip="Vertical frequency scaling factor.", + ) + .req_float_horizontal_factor( + default=1.0, + tooltip="Horizontal frequency scaling factor.", + ) + .req_bool_use_sqrt( + default=True, + tooltip="Controls whether to sqrt when dividing the FFT. Negative hfac/wfac won't work when enabled. Turning it off seems to make the parameters have a much stronger effect.", + ), + ) @classmethod def get_item_class(cls): @@ -180,55 +127,36 @@ class SonarAdvanced1fNoiseNode(SonarCustomNoiseNodeBase): class SonarAdvancedPowerLawNoiseNode(SonarCustomNoiseNodeBase): - DESCRIPTION = "Custom noise type that allows specifying parameters for power law (grey, violet, etc) variants. " + DESCRIPTION = "Custom noise type that allows specifying parameters for power law (grey, violet, etc) variants." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "alpha": ( - "FLOAT", - { - "default": 0.5, - "step": 0.001, - "min": -1000.0, - "max": 1000.0, - "round": False, - "tooltip": "Alpha parameter of the generated noise. Positive values (low frequency noise) tend to produce colorful results.", - }, + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_float_alpha( + default=0.5, + tooltip="Similar to the advanced power noise node, positive values increase low frequencies (with colorful effects), negative values increase high frequencies.", + ) + .req_field_div_max_dims( + ( + "none", + "non-batch", + "spatial", + "all", + "batch", + "channel", + "height", + "width", ), - "div_max_dims": ( - ( - "none", - "non-batch", - "spatial", - "all", - "batch", - "channel", - "height", - "width", - ), - { - "default": "non-batch", - "tooltip": "If non-none, the noise gets divide by the maxmimu over this dimension.", - }, - ), - "use_div_max_abs": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Only has an effect when div_max_dims is not none. Controls whether maximization is done with the absolute values or raw values.", - }, - ), - "use_sign": ( - "BOOLEAN", - { - "default": False, - "tooltip": "When set, only the sign of the initial noise is used, so -0.5, -0.2 all turn into -1, 0.5, 2, etc all turn into 1.", - }, - ), - } - return result + default="non-batch", + tooltip="If non-none, the noise gets divide by the maximum over this dimension.", + ) + .req_bool_use_div_max_abs( + default=True, + tooltip="Only has an effect when div_max_dims is not none. Controls whether maximization is done with the absolute values or raw values.", + ) + .req_bool_use_sign( + tooltip="When set, only the sign of the initial noise is used, so -0.5, -0.2 all turn into -1, 0.5, 2, etc all turn into 1.", + ), + ) @classmethod def get_item_class(cls): @@ -270,200 +198,117 @@ class SonarAdvancedPowerLawNoiseNode(SonarCustomNoiseNodeBase): class SonarAdvancedCollatzNoiseNode(SonarCustomNoiseNodeBase): DESCRIPTION = "Custom noise type that allows specifying parameters for Collatz noise. Very experimental, also very slow. It might just about work as initial noise with non-ancestral sampling but if you get weird results I recommend mixing it with other noise types or possibly using ancestral/SDE sampling." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "adjust_scale": ( - "BOOLEAN", - { - "default": False, - "tooltip": "When enabled, the output will be normalized to values between -1 and 1 using the last two dimensions (if there are four or more), otherwise dimensions after the first.", - }, + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_bool_adjust_scale( + default=False, + tooltip="When enabled, the output will be normalized to values between -1 and 1 using the last two dimensions (if there are four or more), otherwise dimensions after the first.", + ) + .req_string_chain_length( + default="1, 1, 2, 2, 3, 3", + tooltip="Comma-separated list of chain lengths. Cannot be empty. Iterations will cycle through the list and wrap. Controls the length of Collatz chains. Note: Using a high chain length may be very slow, especially if combined with many iterations.", + ) + .req_int_chain_offset( + default=5, + min=0, + max=10000, + tooltip="Uses values starting at the specified offset. Note: This entails generating chains of length chain_length + chain_offset, which may be quite slow if you use high values.", + ) + .req_int_iterations( + default=10, + min=1, + max=10000, + tooltip="Number of iterations to run. Warning: Collatz noise (my implementation, anyway) is EXTREMELY slow.", + ) + .req_bool_iteration_sign_flipping( + default=True, + tooltip="Controls whether we cycle between flipping the sign on the output from each iteration. May average out weirdness... Or make stuff weirder.", + ) + .req_float_rmin( + default=-8000.0, + tooltip="Minimum value a chain can start with. Going as low as -9500 should be safe with float32.", + ) + .req_float_rmax( + default=8000.0, + tooltip="Maximum value a chain can start with. I don't recommend going over 9500 if you are using the float32 dtype here as that is where the Collatz chain starts to reach values that can't be accurately represented.", + ) + .req_string_dims( + default="-1, -1, -2, -2", + tooltip="Comma-separated list of dimensions. Cannot be empty. May be negative to count from the end of the list. Iterations will cycle through the list and wrap.", + ) + .req_bool_flatten( + tooltip="Controls whether dimensions past the current one selected from the dims parameter will get flattened.", + ) + .req_field_output_mode( + ( + "values", + "ratios", + "mults", + "adds", + "seed_x_mults", + "seed_x_adds", + "noise_x_ratios", + "noise_x_mults", + "noise_x_adds", ), - "chain_length": ( - "STRING", - { - "default": "1, 1, 2, 2, 3, 3", - "tooltip": "Comma-separated list of chain lengths. Cannot be empty. Iterations will cycle through the list and wrap. Controls the length of Collatz chains. Note: Using a high chain length may be very slow, especially if combined with many iterations.", - }, - ), - "chain_offset": ( - "INT", - { - "default": 5, - "min": 0, - "max": 10000, - "tooltip": "Uses values starting at the specified offset. Note: This entails generating chains of length chain_length + chain_offset, which may be quite slow if you use high values.", - }, - ), - "iterations": ( - "INT", - { - "default": 10, - "min": 1, - "max": 10000, - "tooltip": "Number of iterations to run. Warning: Collatz noise (my implementation, anyway) is EXTREMELY slow.", - }, - ), - "iteration_sign_flipping": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether we cycle between flipping the sign on the output from each iteration. May average out weirdness... Or make stuff weirder.", - }, - ), - "rmin": ( - "FLOAT", - { - "default": -8000.0, - "min": -100000.0, - "max": 100000.0, - "tooltip": "Minimum value a chain can start with. Going as low as -9500 should be safe with float32.", - }, - ), - "rmax": ( - "FLOAT", - { - "default": 8000.0, - "min": -100000.0, - "max": 100000.0, - "tooltip": "Maximum value a chain can start with. I don't recommend going over 9500 if you are using the float32 dtype here as that is where the Collatz chain starts to reach values that can't be accurately represented.", - }, - ), - "dims": ( - "STRING", - { - "default": "-1, -1, -2, -2", - "tooltip": "Comma-separated list of dimensions. Cannot be empty. May be negative to count from the end of the list. Iterations will cycle through the list and wrap.", - }, - ), - "flatten": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether dimensions past the current one selected from the dims parameter will get flattened.", - }, - ), - "output_mode": ( - ( - "values", - "ratios", - "mults", - "adds", - "seed_x_mults", - "seed_x_adds", - "noise_x_ratios", - "noise_x_mults", - "noise_x_adds", - ), - { - "default": "values", - }, - ), - "quantile": ( - "FLOAT", - { - "default": 0.5, - "min": 0.0, - "max": 1.0, - "tooltip": "The initial output of each iteration will be run through quantile normalization. Setting the parameter to 0 or 1 will disable quantile normalization.", - }, - ), - "quantile_strategy": ( - tuple(utils.quantile_handlers.keys()), - { - "default": "clamp", - "tooltip": "Determines how to treat outliers. zero and reverse_zero modes are only useful if you're going to do something like add the result to some other noise. zero will return zero for anything outside the quantile range, reverse_zero only _keeps_ the outliers and zeros everything else.", - }, - ), - "noise_dtype": ( - ("float32", "float64", "float16", "bfloat16"), - { - "default": "float32", - "tooltip": "Generally should be left at the default. Only float32 and float64 will work if you have quantile normalization enabled.", - }, - ), - "even_multiplier": ( - "FLOAT", - { - "default": 0.5, - "min": -10000.0, - "max": 1000.0, - "tooltip": "Multiplier to use when the previous link in the chain is even. Collatz uses 0.5 (divides by two) here.", - }, - ), - "even_addition": ( - "FLOAT", - { - "default": 0.0, - "min": -10000.0, - "max": 1000.0, - "tooltip": "Value to add when the previous link in the chain is even. Collatz uses 0 here.", - }, - ), - "odd_multiplier": ( - "FLOAT", - { - "default": 3.0, - "min": -10000.0, - "max": 1000.0, - "tooltip": "Multiplier to use when the previous link in the chain is odd. Collatz uses 3 here.", - }, - ), - "odd_addition": ( - "FLOAT", - { - "default": 1.0, - "min": -10000.0, - "max": 1000.0, - "tooltip": "Value to add when the previous link in the chain is odd. Collatz uses 1 here.", - }, - ), - "integer_math": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether the results during chain generation get truncated to an integer value or not. Should be enabled if you actually want to generate accurate Collatz chains.", - }, - ), - "add_preserves_sign": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether additions use the same sign as the item they're being added to.", - }, - ), - "break_loops": ( - "BOOLEAN", - { - "default": True, - "tooltip": "Controls whether the chain resets back to the seed value once it reaches 1 or 0. Generally should be left enabled, otherwise the chain will oscillate between only a few values for the rest of the length (at least with the Collatz rules).", - }, - ), - "seed_mode": ( - ("default", "force_odd", "force_even"), - { - "default": "default", - "tooltip": "Default mode just uses whatever the original seed value was. force_odd/force_even will force it to the specified parity by adding one if it doesn't match. Starting from odd seeds might result in longer chains. Enabling the force modes may cause the initial seeds to exceed rmax by one.", - }, - ), - } - result["optional"] |= { - "seed_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional custom noise to use for initial values for Collatz chains. May be slow as it will generate noise according to the original input size and then crop it. Does this noise type have enough warnings about it being slow? Yeah. Connecting something here will probably make it even slower!\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - "mix_custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional custom noise to use with the output modes starting with 'noise'.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - } - return result + default="values", + ) + .req_float_quantile( + default=0.5, + min=0.0, + max=1.0, + tooltip="The initial output of each iteration will be run through quantile normalization. Setting the parameter to 0 or 1 will disable quantile normalization.", + ) + .req_field_quantile_strategy( + tuple(utils.quantile_handlers.keys()), + default="clamp", + tooltip="Determines how to treat outliers. zero and reverse_zero modes are only useful if you're going to do something like add the result to some other noise. zero will return zero for anything outside the quantile range, reverse_zero only _keeps_ the outliers and zeros everything else.", + ) + .req_field_noise_dtype( + ("float32", "float64", "float16", "bfloat16"), + default="float32", + tooltip="Generally should be left at the default. Only float32 and float64 will work if you have quantile normalization enabled.", + ) + .req_float_even_multiplier( + default=0.5, + tooltip="Multiplier to use when the previous link in the chain is even. Collatz uses 0.5 (divides by two) here.", + ) + .req_float_even_addition( + default=0.0, + tooltip="Value to add when the previous link in the chain is even. Collatz uses 0 here.", + ) + .req_float_odd_multiplier( + default=3.0, + tooltip="Multiplier to use when the previous link in the chain is odd. Collatz uses 3 here.", + ) + .req_float_odd_addition( + default=1.0, + tooltip="Value to add when the previous link in the chain is odd. Collatz uses 1 here.", + ) + .req_bool_integer_math( + default=True, + tooltip="Controls whether the results during chain generation get truncated to an integer value or not. Should be enabled if you actually want to generate accurate Collatz chains.", + ) + .req_bool_add_preserves_sign( + default=True, + tooltip="Controls whether additions use the same sign as the item they're being added to.", + ) + .req_bool_break_loops( + default=True, + tooltip="Controls whether the chain resets back to the seed value once it reaches 1 or 0. Generally should be left enabled, otherwise the chain will oscillate between only a few values for the rest of the length (at least with the Collatz rules).", + ) + .req_field_seed_mode( + ("default", "force_odd", "force_even"), + default="default", + tooltip="Default mode just uses whatever the original seed value was. force_odd/force_even will force it to the specified parity by adding one if it doesn't match. Starting from odd seeds might result in longer chains. Enabling the force modes may cause the initial seeds to exceed rmax by one.", + ) + .opt_customnoise_seed_custom_noise( + tooltip="Optional custom noise to use for initial values for Collatz chains. May be slow as it will generate noise according to the original input size and then crop it. Does this noise type have enough warnings about it being slow? Yeah. Connecting something here will probably make it even slower!", + ) + .opt_customnoise_mix_custom_noise( + tooltip="Optional custom noise to use with the output modes starting with 'noise'.", + ), + ) @classmethod def get_item_class(cls): @@ -640,133 +485,71 @@ class SonarWaveletNoiseNode( ): DESCRIPTION = "Custom noise type that allows generating wavelet noise. Very simple explanation of how a single octave works:\n1) Generate some noise.\n2) Scale it down 50%.\n3) Scale it back up to the original size.\n4) Subtract the scaled noise from the original noise.\nScaling the noise down and then back up blurs it, so this is essentially sharpening the noise." - @classmethod - def INPUT_TYPES(cls): - result = super().INPUT_TYPES() - result["required"] |= { - "octaves": ( - "INT", - { - "default": 4, - "min": -100, - "max": 100, - "tooltip": "Number of octaves to generate. You can use a negative number here to run the octaves in reverse order though it may produce weird results/not work very well.", - }, - ), - "octave_height_factor": ( - "FLOAT", - { - "default": 0.5, - "min": 0.001, - "max": 10000.0, - "tooltip": "Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.", - }, - ), - "octave_width_factor": ( - "FLOAT", - { - "default": 0.5, - "min": 0.001, - "max": 10000.0, - "tooltip": "Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.", - }, - ), - "octave_scale_mode": ( - utils.UPSCALE_METHODS, - { - "tooltip": "Scaling mode used within each octave to produce the scaled noise. By default this will be scaling down that octave's noise.", - "default": "adaptive_avg_pool2d", - }, - ), - "octave_rescale_mode": ( - utils.UPSCALE_METHODS, - { - "tooltip": "Scaling mode used within each octave to scale the noise back up to that octave's original size.", - "default": "bilinear", - }, - ), - "post_octave_rescale_mode": ( - utils.UPSCALE_METHODS, - { - "tooltip": "Scaling mode used to scale the output of an octave back up to the actual latent size.", - "default": "bilinear", - }, - ), - "initial_amplitude": ( - "FLOAT", - { - "default": 1.0, - "min": -10000.0, - "max": 10000.0, - "tooltip": "Basically the strength an octave gets added to the total. This will be scaled by persistance after each octave.", - }, - ), - "persistence": ( - "FLOAT", - { - "default": 0.5, - "min": -10000.0, - "max": 10000.0, - "tooltip": "Multiplier applied to amplitude after each octave. 0.5 means the first octave uses initial_amplitude, the second uses half of that and so on.", - }, - ), - "height_factor": ( - "FLOAT", - { - "default": 2.0, - "min": 0.001, - "max": 10000.0, - "tooltip": "Scaling factor for height, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.", - }, - ), - "width_factor": ( - "FLOAT", - { - "tooltip": "Scaling factor for width, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.", - "default": 2.0, - "min": 0.001, - "max": 10000.0, - }, - ), - "update_blend": ( - "FLOAT", - { - "tooltip": "Controls how original_noise - scaled_noise is blended with original_noise. The default is to use 100% original_noise - scaled_noise.", - "default": 1.0, - "min": -10000.0, - "max": 10000.0, - }, - ), - "update_blend_mode": ( - ("simple_add", *utils.BLENDING_MODES.keys()), - { - "default": "lerp", - "tooltip": "Controls how the enhanced noise from each octave is blended with that octave's raw noise. With normal wavelet noise there's no blending and you use 100% enhanced noise.", - }, - ), - "normalize_noise": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Controls whether the noise source is normalized before wavelet filtering occurs.", - }, - ), - "normalize": ( - ("default", "forced", "disabled"), - { - "tooltip": "Controls whether the generated noise is normalized to 1.0 strength. For weird blend modes, you may want to set this to forced.", - }, - ), - } - result["optional"] |= { - "custom_noise": ( - WILDCARD_NOISE, - { - "tooltip": f"Optional: Custom noise input. If unconnected will default to Gaussian noise. Note: When connected, the noise for all octaves will be generated at the maximum scale and then cropped which may be slow.\n{NOISE_INPUT_TYPES_HINT}", - }, - ), - } - return result + INPUT_TYPES = SonarLazyInputTypes( + lambda: NoiseChainInputTypes() + .req_int_octaves( + default=4, + min=-100, + max=100, + tooltip="Number of octaves to generate. You can use a negative number here to run the octaves in reverse order though it may produce weird results/not work very well.", + ) + .req_float_octave_height_factor( + default=0.5, + min=0.001, + tooltip="Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.", + ) + .req_float_octave_width_factor( + default=0.5, + min=0.001, + tooltip="Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.", + ) + .req_selectscalemode_octave_scale_mode( + default="adaptive_avg_pool2d", + tooltip="Scaling mode used within each octave to produce the scaled noise. By default this will be scaling down that octave's noise.", + ) + .req_selectscalemode_octave_rescale_mode( + default="bilinear", + tooltip="Scaling mode used within each octave to scale the noise back up to that octave's original size.", + ) + .req_selectscalemode_post_octave_rescale_mode( + default="bilinear", + tooltip="Scaling mode used to scale the output of an octave back up to the actual latent size.", + ) + .req_float_initial_amplitude( + default=1.0, + tooltip="Basically the strength an octave gets added to the total. This will be scaled by persistance after each octave.", + ) + .req_float_persistence( + default=0.5, + tooltip="Multiplier applied to amplitude after each octave. 0.5 means the first octave uses initial_amplitude, the second uses half of that and so on.", + ) + .req_float_height_factor( + default=2.0, + min=0.001, + tooltip="Scaling factor for height, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.", + ) + .req_float_width_factor( + tooltip="Scaling factor for width, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.", + default=2.0, + min=0.001, + ) + .req_float_update_blend( + tooltip="Controls how original_noise - scaled_noise is blended with original_noise. The default is to use 100% original_noise - scaled_noise.", + default=1.0, + ) + .req_selectblend_update_blend_mode( + insert_modes=("simple_add",), + default="lerp", + tooltip="Controls how the enhanced noise from each octave is blended with that octave's raw noise. With normal wavelet noise there's no blending and you use 100% enhanced noise.", + ) + .req_bool_normalize_noise( + tooltip="Controls whether the noise source is normalized before wavelet filtering occurs.", + ) + .req_normalizetristate_normalize() + .opt_customnoise_custom_noise( + tooltip="Optional: Custom noise input. If unconnected will default to Gaussian noise. Note: When connected, the noise for all octaves will be generated at the maximum scale and then cropped which may be slow.", + ), + ) @classmethod def get_item_class(cls): @@ -820,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/powernoise.py b/py/nodes/powernoise.py similarity index 79% rename from py/powernoise.py rename to py/nodes/powernoise.py index d791ee8..d4bb137 100644 --- a/py/powernoise.py +++ b/py/nodes/powernoise.py @@ -16,14 +16,16 @@ from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler from PIL import Image from torch import Tensor -from .nodes.base 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 PREVIEW_FORMAT = comfy.latent_formats.SD15() @@ -555,122 +557,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): @@ -695,7 +644,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", @@ -876,67 +825,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( diff --git a/py/noise.py b/py/noise.py index 7feeaaf..491b1b2 100644 --- a/py/noise.py +++ b/py/noise.py @@ -2,6 +2,7 @@ from __future__ import annotations import abc import math +import random from functools import partial from typing import Callable @@ -15,6 +16,7 @@ from . import external, utils from .noise_generation import * from .sonar import SonarGuidanceMixin from .utils import ( + RNGStates, crop_samples, fallback, pattern_break, @@ -40,7 +42,15 @@ class CustomNoiseItemBase(abc.ABC): self.factor = factor self.keys = set(kwargs.keys()) for k, v in kwargs.items(): - setattr(self, k, v) + do_clone = k in { + "custom_noise", + "custom_noise_opt", + "noise", + "noise_opt", + "sonar_custom_noise", + "sonar_custom_noise_opt", + } and hasattr(v, "clone") + setattr(self, k, v.clone() if do_clone else v) def clone_key(self, k): return getattr(self, k) @@ -433,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, @@ -1179,9 +1213,7 @@ class NormalizeToScaleNoise(CustomNoiseItemBase): min_positive_value: float, max_positive_value: float, mode: str, - dims: tuple, - normalize_noise: float, - normalize, + **kwargs, ): if mode == "simple": if min_negative_value >= max_positive_value: @@ -1207,9 +1239,7 @@ class NormalizeToScaleNoise(CustomNoiseItemBase): min_positive_value=min_positive_value, max_positive_value=max_positive_value, mode=mode, - dims=dims, - normalize_noise=normalize_noise, - normalize=normalize, + **kwargs, ) def clone_key(self, k): @@ -1218,6 +1248,8 @@ class NormalizeToScaleNoise(CustomNoiseItemBase): return super().clone_key(k) def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + std_dims, std_multiplier = self.std_dims, self.std_multiplier + mean_dims, mean_multiplier = self.mean_dims, self.mean_multiplier factor = self.factor mode = self.mode if mode == "simple": @@ -1252,6 +1284,16 @@ class NormalizeToScaleNoise(CustomNoiseItemBase): else: for bidx in range(noise.shape[0]): noise[bidx] = noise_filter(noise[bidx]) + if mean_multiplier != 0: + noise -= noise.mean(dim=mean_dims, keepdim=True).mul_(mean_multiplier) + if std_multiplier != 0: + noise_std = ( + noise.std(dim=std_dims, keepdim=True) + .sub_(1.0) + .mul_(std_multiplier) + .add_(1.0) + ) + noise /= torch.where(noise_std == 0, 1e-07, noise_std) return scale_noise(noise, factor, normalized=normalize) return noise_sampler @@ -1266,26 +1308,37 @@ 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: + 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__( factor, noise_2_percent=noise_2_percent, 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, ) @@ -1294,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): @@ -1301,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, @@ -1318,13 +1377,30 @@ 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 n2_blend == 0 or ns_2 is None - else blend_function(noise_1, ns_2(s, sn), n2_blend_tensor) + if noise_2 is None + else blend_function(noise_1, noise_2, n2_blend_tensor) ) return scale_noise(noise, factor, normalized=normalize) @@ -1353,9 +1429,23 @@ class ResizedNoise(CustomNoiseItemBase): raise ValueError("ResizedNoise can only handle 3+ dimensional latents") factor = self.factor normalize = self.get_normalize("normalize", normalized) + spatial_compression = self.spatial_compression + spatial_mode = self.spatial_mode + width, height = self.width, self.height xh, xw = x.shape[-2:] - nh, nw = self.height // 8, self.width // 8 - offsh, offsw = self.crop_offset_vertical // 8, self.crop_offset_horizontal // 8 + if spatial_mode != "percentage": + height //= spatial_compression + width //= spatial_compression + if spatial_mode == "absolute": + nh, nw = int(height), int(width) + elif spatial_mode == "relative": + nh, nw = int(xh + height), int(xw + width) + elif spatial_mode == "percentage": + nh, nw = max(1, int(xh * height)), max(1, int(xw * width)) + else: + raise ValueError("Bad spatial_mode") + offsh = self.crop_offset_vertical // spatial_compression + offsw = self.crop_offset_horizontal // spatial_compression if xh == nh and xw == nw: ns = self.custom_noise.make_noise_sampler( x, @@ -1939,6 +2029,116 @@ class PatternBreakNoise(CustomNoiseItemBase): return noise_sampler +class CustomNoiseParametersNoise(CustomNoiseItemBase): + def clone_key(self, k): + if k == "noise": + return self.noise.clone() + return super().clone_key(k) + + def make_noise_sampler( + self, + x, + sigma_min, + sigma_max, + *args, + normalized=True, + **kwargs, + ): + factor = self.factor + normalize = self.get_normalize("normalize", normalized) + orig_shape = x.shape + orig_dtype = x.dtype + orig_device = x.device + if self.override_device is not None: + kwargs["cpu"] = self.override_device == "cpu" + x = x.to(device=self.override_device) + if x.ndim == 5 and self.frames_to_channels: + x = x.reshape(x.shape[0], x.shape[1] * x.shape[2], *x.shape[3:]) + fix_invalid = self.fix_invalid + if self.override_dtype and x.dtype != self.override_dtype: + x = x.to(dtype=self.override_dtype) + fixed_aspect = False + if self.ensure_square_aspect_ratio: + if x.ndim == 3: + height, width = 1, x.shape[-1] + spatdims = 1 + else: + spatdims = 2 + height, width = x.shape[-2:] + hw = (height * width) ** 0.5 + if not hw.is_integer(): + fixed_aspect = True + hw = math.ceil(hw) + temp_x = x.new_zeros(*x.shape[:-spatdims], hw**2) + temp_x[..., : height * width] = x.flatten(start_dim=-spatdims)[ + ..., + : height * width, + ] + x = temp_x.reshape(*temp_x.shape[:-1], hw, hw) + if self.rng_offset_mode in {"override", "add"}: + seed = ( + self.rng_state_offset + if self.rng_offset_mode == "override" + else kwargs.pop("seed", 0) + self.rng_state_offset + ) + kwargs["seed"] = seed + else: + seed = kwargs.get("seed", 0) + rng_mode = self.rng_mode + if rng_mode == "separate": + rng_state = RNGStates(x.device.type) + if self.rng_offset_mode != "disabled": + temp_rng_state = rng_state + try: + random.seed(seed) + torch.manual_seed(seed) + rng_state = RNGStates(x.device.type) + finally: + temp_rng_state.set_states() + del temp_rng_state + else: + rng_state = None + ns = self.noise.make_noise_sampler( + x, + *args, + sigma_min=sigma_min, + sigma_max=sigma_max, + normalized=False, + **kwargs, + ) + device_type = x.device.type + + def noise_sampler(sigma, sigma_next) -> torch.Tensor: + if rng_mode != "default": + temp_rng_state = RNGStates(device_type) + try: + if rng_mode == "separate": + rng_state.set_states() + noise = ns(sigma, sigma_next) + if rng_mode == "separate": + rng_state.update() + finally: + temp_rng_state.set_states() + else: + noise = ns(sigma, sigma_next) + if fix_invalid: + noise_temp = noise.nan_to_num(0, posinf=0, neginf=0) + noise = noise.nan_to_num_( + 0, + posinf=noise_temp.max(), + neginf=noise_temp.min(), + ) + if fixed_aspect: + noise = noise.flatten(start_dim=-spatdims)[..., : height * width] + if noise.shape != orig_shape: + noise = noise.reshape(orig_shape) + if noise.dtype != orig_dtype or noise.device != orig_device: + noise = noise.to(device=orig_device, dtype=orig_dtype) + return scale_noise(noise, factor, normalized=normalize) + + return noise_sampler + + class BlehOpsNoise(CustomNoiseItemBase): def __init__( self, @@ -2169,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 0bebff7..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" @@ -1424,10 +1885,12 @@ class ScatternetFilteredNoiseGenerator(FramesToChannelsNoiseGenerator): ) super().__init__(*args, **kwargs) if self.output_mode not in { - "channels_adjusted", "channels", + "channels_adjusted", + "channels_scaled", "flat", "flat_adjusted", + "flat_scaled", }: raise ValueError("Bad output mode") @@ -1462,6 +1925,8 @@ class ScatternetFilteredNoiseGenerator(FramesToChannelsNoiseGenerator): "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, } @@ -1479,12 +1944,14 @@ class ScatternetFilteredNoiseGenerator(FramesToChannelsNoiseGenerator): def generate(self, *args): adjusted_shape = self.get_adjusted_shape() - adjusted = self.output_mode.endswith("_adjusted") + 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 2 ** abs(self.scatternet_order) + spatial_compensation = 1 if adjusted else order_spatial_compensation if self.noise_sampler is None: temp_shape = ( ( @@ -1498,6 +1965,20 @@ class ScatternetFilteredNoiseGenerator(FramesToChannelsNoiseGenerator): 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) @@ -1729,7 +2210,7 @@ class CollatzNoiseGenerator(NoiseGenerator): result[dim] = slice(idx, None, stride) return result - def _generate_iteration( # noqa: PLR0914 + def _generate_iteration( self, *args, dim: int, @@ -1996,6 +2477,7 @@ __all__ = ( "ScatternetFilteredNoiseGenerator", "StudentTNoiseGenerator", "UniformNoiseGenerator", + "VoronoiNoiseGenerator", "WaveletFilteredNoiseGenerator", "WaveletNoiseGenerator", ) diff --git a/py/utils.py b/py/utils.py index 2f33fc0..89f9b51 100644 --- a/py/utils.py +++ b/py/utils.py @@ -1,14 +1,19 @@ from __future__ import annotations import math +import random from functools import partial +from typing import TYPE_CHECKING, Callable import torch -from comfy.model_management import device_supports_non_blocking +from comfy.model_management import device_supports_non_blocking, get_torch_device from comfy.utils import common_upscale from .external import MODULES as EXT +if TYPE_CHECKING: + from collections.abc import Sequence + BLENDING_MODES = { "lerp": torch.lerp, "inject": lambda a, b, t: (b * t).add_(a), @@ -25,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, @@ -99,8 +129,12 @@ def _quantile_norm_scaledown( **_kwargs: dict, ) -> torch.Tensor: noiseabs = noise.abs() - mv = noiseabs.max(dim=dim, keepdim=True).clamp(min=1e-06) - return noise if mv == 0 else torch.where(noiseabs > nq, noise * (nq / mv), noise) + mv = noiseabs.max(dim=dim, keepdim=True).values.clamp(min=1e-06) + return ( + noise + if mv.sum().item() == 0 + else torch.where(noiseabs > nq, noise * (nq / mv), noise) + ) def _quantile_norm_wave( @@ -141,6 +175,24 @@ def _quantile_norm_mode( ) +def _quantile_norm_replace( + noise: torch.Tensor, + nq: torch.Tensor, + *, + keep_sign: bool = False, + avoid_sign: bool = False, + **_kwargs: dict, +) -> torch.Tensor: + mask = noise.abs() <= nq + candidates = noise[mask].flatten() + candidates = candidates[torch.arange(noise.numel()) % candidates.numel()].reshape( + noise.shape, + ) + if keep_sign or avoid_sign: + candidates = candidates.copysign_(noise.neg() if avoid_sign else noise) + return torch.where(mask, noise, candidates) + + quantile_handlers = { "clamp": lambda noise, nq, **_kwargs: noise.clamp(-nq, nq), "scale_down": _quantile_norm_scaledown, @@ -235,6 +287,9 @@ quantile_handlers = { ), "mode_1dec": partial(_quantile_norm_mode, decimals=1), "mode_2dec": partial(_quantile_norm_mode, decimals=2), + "replace": _quantile_norm_replace, + "replace_keepsign": partial(_quantile_norm_replace, keep_sign=True), + "replace_avoidsign": partial(_quantile_norm_replace, avoid_sign=True), } @@ -484,3 +539,126 @@ def trunc_decimals(x: torch.Tensor, decimals: int = 3) -> torch.Tensor: def maybe_apply(val, cond, fun): return fun(val) if cond else val + + +def maybe_apply_kwargs(d: dict | None, cond, fun, *, default=None): + return default if d is None or not cond else fun(**d) + + +def tensor_item(val: torch.Tensor | float, *, collapse_function=torch.max) -> float: + if isinstance(val, torch.Tensor): + return float(collapse_function(val).detach().cpu().item()) + return float(val) + + +# Does not handle out of order or duplicated sigmas. +def step_from_sigmas( + sigma: float | torch.Tensor, + sigmas: torch.Tensor, + *, + decimals: int | None = 4, + output_decimals: int = 2, +) -> float | None: + sigma = tensor_item(sigma) + sigmas = sigmas.detach().cpu() + if sigmas.ndim == 2: + sigmas = sigmas.max(dim=0).values + elif sigmas.ndim != 1: + errstr = f"Unexpected number of dimensions in sigmas, should be 1 or 2 but got shape {sigmas.shape}" + raise ValueError(errstr) + sigmas = sigmas[:-1] + if not len(sigmas) or torch.any(sigmas <= 0): + return None + if decimals is not None: + sigmas = sigmas.round(decimals=decimals) + sigma = round(sigma, decimals) + sigma_min, sigma_max = sigmas.aminmax() + if not sigma_min <= sigma <= sigma_max: + return None + max_idx = len(sigmas) - 1 + idx = int(tensor_item((sigmas - sigma).abs().argmin())) + idx_sigma = tensor_item(sigmas[idx]) + if decimals is not None: + idx_sigma = round(idx_sigma, decimals) + if sigma == idx_sigma: + return float(idx) + # Between sigmas, but guaranteed to be in range here. + idx_low, idx_high = (idx, idx - 1) if sigma > idx_sigma else (idx + 1, idx) + if idx_low < 0 or idx_high < 0 or idx_low > max_idx or idx_high > max_idx: + return None + sigma_low, sigma_high = tensor_item(sigmas[idx_low]), tensor_item(sigmas[idx_high]) + step_diff = sigma_high - sigma_low + if step_diff == 0: + return float(idx) + pct = 1.0 - ((sigma - sigma_low) / step_diff) + return round(idx_high + pct, output_decimals) + + +def clamp_float(val: float, minval=0.0, maxval=1.0) -> float: + return max(minval, min(val, maxval)) + + +def filter_dict(d: dict, keep: set | Sequence, *, recursive: bool = False) -> dict: + return { + k: v if not (recursive and isinstance(v, dict)) else filter_dict(v, keep) + for k, v in d.items() + if k in keep + } + + +class RNGStates: + DEFAULT_GPU_TYPE = get_torch_device().type + + def __init__( + self, + device_types: set | str | Sequence | None = None, + *, + add_defaults: bool = True, + ): + if device_types is None: + device_types = set() + elif isinstance(device_types, str): + device_types = {device_types} + elif not isinstance(device_types, set): + device_types = set(device_types) + if add_defaults: + device_types = device_types | {"python", "cpu", self.DEFAULT_GPU_TYPE} # noqa: PLR6104 + self.rng_states = self.get_states(device_types) + + def update(self): + self.rng_states = self.get_states(set(self.rng_states)) + + @staticmethod + def get_states(device_types: set) -> dict: + return { + k: torch.get_rng_state() + if k == "cpu" + else ( + random.getstate() + if k == "python" + else getattr(torch, k).get_rng_state() + ) + for k in device_types + if k in {"python", "cpu"} or hasattr(torch, k) + } + + def set_states(self, *, update: bool = True, override_states: dict | None = None): + states = self.rng_states if override_states is None else override_states + new_states = {} + for k, v in states.items(): + if isinstance(v, torch.Tensor): + v = v.clone() # noqa: PLW2901 + if k == "cpu": + new_states[k] = v + torch.set_rng_state(v) + continue + if k == "python": + new_states[k] = v + random.setstate(v) + continue + tm = getattr(torch, k, None) + if tm is not None: + new_states[k] = v + tm.set_rng_state(v) + if update: + self.rng_states = new_states diff --git a/py/wavelet_cfg.py b/py/wavelet_cfg.py new file mode 100644 index 0000000..da854bb --- /dev/null +++ b/py/wavelet_cfg.py @@ -0,0 +1,842 @@ +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 + 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: + step_first = enabled_steps[0].item() + step_last = enabled_steps[-1].item() + pct_enabled_steps = (step - step_first) / (step_last - step_first) + else: + 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}" + 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() diff --git a/py/wavelet_functions.py b/py/wavelet_functions.py index 1951184..638f2a0 100644 --- a/py/wavelet_functions.py +++ b/py/wavelet_functions.py @@ -11,17 +11,19 @@ if TYPE_CHECKING: try: import pytorch_wavelets as ptwav + import pywt HAVE_WAVELETS = True except ImportError: ptwav = None + pywt = None HAVE_WAVELETS = False class Wavelet: - DEFAULT_MODE = "periodization" + DEFAULT_MODE = "symmetric" DEFAULT_LEVEL = 3 - DEFAULT_WAVE = "haar" + DEFAULT_WAVE = "db4" DEFAULT_USE_1D_DWT = False DEFAULT_USE_DTCWT = False DEFAULT_QSHIFT = "qshift_a" @@ -102,16 +104,97 @@ class Wavelet: )) return result - def to(self, *args: list, **kwargs: dict) -> None: - self._wavelet_forward = self._wavelet_forward.to(*args, **kwargs) - self._wavelet_inverse = self._wavelet_inverse.to(*args, **kwargs) + def to(self, *args: list, copy: bool = False, **kwargs: dict) -> 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 | torch.Tensor | None, + yh_scales: float | Sequence | None, *, in_place: bool = False, ) -> tuple: @@ -120,19 +203,16 @@ def wavelet_scaling( yh = tuple(yhband.clone() for yhband in yh) if yl_scale != 1.0: yl *= yl_scale - if yh_scales is None or yh_scales == 1.0: - return (yl, yh) - if isinstance(yh_scales, (int, float)): - yh_scales = (yh_scales,) * len(yh) - # print("SCALES", self.yl_scale, yh_scales) + 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): - # print(">> SCALING", hscale) if isinstance(hscale, (int, float)): ht *= hscale # noqa: PLW2901 continue for lidx in range(min(ht.shape[2], len(hscale))): - # print(">> SCALE IDX", lidx) - ht[:, :, lidx, :, :] *= hscale[lidx] + ht[:, :, lidx] *= hscale[lidx] return (yl, yh) diff --git a/ruff.toml b/ruff.toml index e589ee0..9c90993 100644 --- a/ruff.toml +++ b/ruff.toml @@ -29,10 +29,12 @@ ignore = [ "FBT002", "PLR0912", "PLR0913", + "PLR0914", "PLR0915", "PLR0917", "PLR2004", "T201", + "TID252", "TRY003", "N802", "N999",