diff --git a/changelog.md b/changelog.md index a65e2c9..8b7a6b7 100644 --- a/changelog.md +++ b/changelog.md @@ -2,12 +2,13 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. -## 20250723 +## 20250727 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. ## 20250705 diff --git a/py/latent_ops.py b/py/latent_ops.py index 03b9c24..53553ed 100644 --- a/py/latent_ops.py +++ b/py/latent_ops.py @@ -14,7 +14,6 @@ if TYPE_CHECKING: class SonarLatentOperation: EXTENDED_LATENT_OPERATION = True - SKIP_ARGS = frozenset(("sigma", "t2", "cond", "uncond", "cond_scale", "raw_args")) def __init__( self, @@ -28,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( @@ -42,8 +43,8 @@ class SonarLatentOperation: if op is None: return t if not getattr(op, "EXTENDED_LATENT_OPERATION", False): - kwargs = {k: v for k, v in kwargs.items() if k not in self.SKIP_ARGS} - return op(t, *args, **kwargs) + return op(latent=t) + return op(*args, latent=t, **kwargs) def __call__( self, diff --git a/py/nodes/base.py b/py/nodes/base.py index 8456a25..d13fa87 100644 --- a/py/nodes/base.py +++ b/py/nodes/base.py @@ -52,6 +52,7 @@ class SonarInputCollection(InputCollection): super().__init__(*args, **kwargs) self._DELEGATE_KEYS = self._DELEGATE_KEYS | frozenset(( # noqa: PLR6104 "customnoise", + "floatpct", "normalizetristate", "selectblend", "selectnoise", @@ -165,6 +166,9 @@ class SonarInputCollection(InputCollection): **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 @@ -180,10 +184,10 @@ class SonarInputTypes(InputTypes): class SonarLazyInputTypes(LazyInputTypes): _NO_REPLACE = True - def __init__(self, *args: list, initializers=(), **kwargs: dict): + def __init__(self, *args: list, initializers=(MODULES.initialize,), **kwargs: dict): super().__init__( *args, - initializers=(MODULES.initialize, *initializers), + initializers=initializers, **kwargs, ) @@ -215,6 +219,7 @@ class SonarCustomNoiseNodeBase(metaclass=IntegratedNode): _skip=not include_chain, tooltip="Optional input for more custom noise items.", ), + initializers=(), ) def go( diff --git a/py/nodes/freeu_extreme.py b/py/nodes/freeu_extreme.py index c9650b2..2e6255d 100644 --- a/py/nodes/freeu_extreme.py +++ b/py/nodes/freeu_extreme.py @@ -3,7 +3,7 @@ from __future__ import annotations import torch from .. import utils -from ..external import IntegratedNode +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 a309b83..a42af7a 100644 --- a/py/nodes/integrations.py +++ b/py/nodes/integrations.py @@ -2,12 +2,12 @@ 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, ) @@ -23,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): @@ -120,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( @@ -146,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): @@ -199,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") @@ -315,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 d2f6eb7..ffc7662 100644 --- a/py/nodes/latent_operations.py +++ b/py/nodes/latent_operations.py @@ -339,20 +339,18 @@ 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. If you attach the optional extra LATENT_OPERATIONS, they will be called in sequence _before_ blending or output scaling." diff --git a/py/nodes/noise_filters.py b/py/nodes/noise_filters.py index 3729e0f..1eba841 100644 --- a/py/nodes/noise_filters.py +++ b/py/nodes/noise_filters.py @@ -1,5 +1,8 @@ from __future__ import annotations +import torch +from comfy import model_management + from .. import noise, utils from ..latent_ops import SonarLatentOperation from ..sonar import SonarGuidanceMixin @@ -1401,24 +1404,147 @@ class SonarLatentOperationFilteredNoiseNode( ) +class SonarCustomNoiseParametersNode( + SonarCustomNoiseNodeBase, + SonarNormalizeNoiseNodeMixin, +): + DESCRIPTION = "Custom noise type that allows overriding some parameters." + + 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, + "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, - "SonarWaveletFilteredNoise": SonarWaveletFilteredNoiseNode, "SonarRippleFilteredNoise": SonarRippleFilteredNoiseNode, - "SonarQuantileFilteredNoise": SonarQuantileFilteredNoiseNode, - "SonarNormalizeNoiseToScale": SonarNormalizeNoiseToScaleNode, - "SonarPerDimNoise": SonarPerDimNoiseNode, "SonarScatternetFilteredNoise": SonarScatternetFilteredNoiseNode, - "SonarLatentOperationFilteredNoise": SonarLatentOperationFilteredNoiseNode, + "SonarScheduledNoise": SonarScheduledNoiseNode, + "SonarShuffledNoise": SonarShuffledNoiseNode, + "SonarWaveletFilteredNoise": SonarWaveletFilteredNoiseNode, } diff --git a/py/nodes/powernoise.py b/py/nodes/powernoise.py index 01c6fd9..d4bb137 100644 --- a/py/nodes/powernoise.py +++ b/py/nodes/powernoise.py @@ -21,7 +21,9 @@ from ..utils import scale_noise from .base import ( NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE, + NoiseChainInputTypes, SonarCustomNoiseNodeBase, + SonarInputTypes, SonarNormalizeNoiseNodeMixin, ) @@ -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/nodes/wavelet_cfg.py b/py/nodes/wavelet_cfg.py index 211a063..09d786d 100644 --- a/py/nodes/wavelet_cfg.py +++ b/py/nodes/wavelet_cfg.py @@ -214,7 +214,7 @@ class WCFGScales(NamedTuple): **_kwargs: dict, ) -> WCFGScales: if verbose: - tqdm.write(f"WCFG: low={self.yl_scale:.4f}, high: {self.yh_scales}") + tqdm.write(f"WCFG: {self.pretty_scales()}") return self def apply_scales( @@ -234,6 +234,22 @@ class WCFGScales(NamedTuple): ) -> 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 @@ -287,12 +303,22 @@ class WCFGScheduledScale(NamedTuple): pct = utils.clamp_float(1.0 - pct) return pct + def pretty_non_default(self) -> str: + result = ", ".join( + f"{fn}={fv.pretty_non_default()}" + if hasattr(fv, "pretty_non_default") + else f"{fn}={fv!r}" + for fn, fv in ((_fn, getattr(self, _fn)) for _fn in self._fields) + if fv != getattr(DEFAULT_SCHEDULEDSCALE, fn) + ) + return f"WCFGScheduledScale({result})" + DEFAULT_SCHEDULEDSCALE = WCFGScheduledScale() class WCFGScalesRange(NamedTuple): - scales_start: WCFGScales + scales_start: WCFGScales = WCFGScales() scales_end: WCFGScales | None = None scheduler: WCFGScheduledScale | None = None @@ -342,7 +368,7 @@ class WCFGScalesRange(NamedTuple): if simple_result is not None: if verbose: tqdm.write( - f"WCFG: low={simple_result.yl_scale:.4f}, high: {simple_result.yh_scales}", + f"WCFG: {simple_result.pretty_scales()}", ) return simple_result start_scale, end_scale = 1.0 - pct, pct @@ -353,11 +379,12 @@ class WCFGScalesRange(NamedTuple): tuple(os * start_scale + oe * end_scale 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: low={yl_scale:.4f}, high: {yh_scales}", + f"WCFG: {result.pretty_scales()}", ) - return WCFGScales(yl_scale=yl_scale, yh_scales=yh_scales) + return result def apply_scales( self, @@ -376,6 +403,19 @@ class WCFGScalesRange(NamedTuple): ) -> tuple[torch.Tensor, Sequence]: return self.get_scales(pcts, yh, verbose=verbose).apply_scales(yl, yh) + def pretty_non_default(self) -> str: + result = ", ".join( + f"{fn}={fv.pretty_non_default()}" + if hasattr(fv, "pretty_non_default") + else f"{fn}={fv!r}" + for fn, fv in ((_fn, getattr(self, _fn)) for _fn in self._fields) + if fv != getattr(DEFAULT_SCALESRANGE, fn) + ) + return f"WCFGScalesRange({result})" + + +DEFAULT_SCALESRANGE = WCFGScalesRange() + class WCFGScheduledFloat(NamedTuple): value_start: float @@ -508,12 +548,22 @@ class WCFGRule(NamedTuple): verbose: bool = False, ) -> tuple[torch.Tensor, Sequence]: scales = getattr(self, name).get_scales(pcts, yh) - if verbose: + if verbose and (scales.yl_scale != 1.0 or scales.yh_scales != 1.0): tqdm.write( - f"WCFG: scales({name:>6}): low={scales.yl_scale:.4f}, high: {scales.yh_scales}", + f"WCFG: scales({name:>6}): {scales.pretty_scales()}", ) return scales.apply_scales(yl, yh) + def pretty_non_default(self) -> str: + result = ", ".join( + f"{fn}={fv.pretty_non_default()}" + if hasattr(fv, "pretty_non_default") + else f"{fn}={fv!r}" + for fn, fv in ((_fn, getattr(self, _fn)) for _fn in self._fields) + if fv != getattr(DEFAULT_RULE, fn) + ) + return f"WCFGRule({result})" + DEFAULT_RULE = WCFGRule() @@ -556,6 +606,7 @@ class WCFGContext(NamedTuple): sigma: torch.Tensor wavelet: Wavelet dtype: torch.dtype + op_kwargs: dict class WaveletCFG: @@ -590,11 +641,22 @@ class WaveletCFG: return x - (cond - uncond).mul_(scale).add_(uncond) @staticmethod - def maybe_op(t: torch.Tensor, mop: Callable | None) -> torch.Tensor: - return t if mop is None else mop(latent=t) + 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 = args["sigma"] + sigma_orig = sigma = args["sigma"] rule_id = id(rule) x = args["input"] if x.ndim == 3 and not rule.use_1d_dwt: @@ -614,8 +676,15 @@ class WaveletCFG: cond, uncond = args["cond_denoised"], args["uncond_denoised"] else: raise ValueError("Bad target mode") - cond = self.maybe_op(cond, self.operation_cond) - uncond = self.maybe_op(uncond, self.operation_uncond) + 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: @@ -635,6 +704,7 @@ class WaveletCFG: sigma=sigma, wavelet=wavelet, dtype=eff_dtype, + op_kwargs=op_kwargs, ) def process_output( @@ -655,7 +725,7 @@ class WaveletCFG: result = ctx.x - result elif rule.target_mode == WCFGTarget.NOISE_NORM: result *= ctx.sigma - return self.maybe_op(result, self.operation_wavelet_cfg) + return self.maybe_op(result, self.operation_wavelet_cfg, **ctx.op_kwargs) @classmethod def wavelet_cfg( @@ -708,7 +778,9 @@ class WaveletCFG: 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}") + 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( @@ -725,6 +797,10 @@ class WaveletCFG: 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) @@ -732,6 +808,7 @@ class WaveletCFG: 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 @@ -739,7 +816,11 @@ class WaveletCFG: 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).contiguous() + return self.maybe_op( + result, + self.operation_result, + **ctx.op_kwargs, + ).contiguous() class SonarWaveletCFGNode(metaclass=IntegratedNode): @@ -973,85 +1054,6 @@ verbose: false return (model,) -# class SonarWaveletCFGSimpleNode(SonarWaveletCFGNode): -# DESCRIPTION = "Wavelet CFG function that allows you to apply different CFG strength to different frequencies (simple version)." - -# INPUT_TYPES = SonarLazyInputTypes( -# lambda: 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.", -# ), -# ) - -# @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, -# ) -> tuple[object]: -# if start_sigma < 0: -# start_sigma = math.inf -# wavelet_params = yaml.safe_load(yaml_parameters) -# rules = WCFGRules.build( -# **( -# { -# "start_sigma": start_sigma, -# "end_sigma": end_sigma, -# "fallback_existing": fallback_mode == "existing", -# "blend_mode": blend_mode, -# "blend_strength": blend_strength, -# } -# | wavelet_params -# ), -# ) -# if len(rules) and rules[0].verbose: -# tqdm.write(f"\nWCFG: Using rules: {rules}\n") -# model = model.clone() -# model.set_model_sampler_cfg_function( -# WaveletCFG( -# existing_cfg=model.model_options.get("sampler_cfg_function"), -# rules=rules, -# operation_cond=operation_cond, -# operation_uncond=operation_uncond, -# operation_fallback_cfg=operation_fallback_cfg, -# operation_wavelet_cfg=operation_wavelet_cfg, -# operation_result=operation_result, -# ), -# ) -# return (model,) - - NODE_CLASS_MAPPINGS = { "SonarWaveletCFG": SonarWaveletCFGNode, } diff --git a/py/noise.py b/py/noise.py index 39429bc..7b4a474 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, @@ -1292,10 +1294,9 @@ class BlendedNoise(CustomNoiseItemBase): raise ValueError( "When custom_noise_2 is not attached noise_2_percent must be set to 0", ) - if noise_2_percent == 1: + if 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, @@ -1337,10 +1338,11 @@ class BlendedNoise(CustomNoiseItemBase): def noise_sampler(s, sn): noise_1 = ns_1(s, sn) + noise_2 = None if ns_2 is None else ns_2(s, sn) 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) @@ -1969,6 +1971,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, diff --git a/py/utils.py b/py/utils.py index 8a27d38..c7b72af 100644 --- a/py/utils.py +++ b/py/utils.py @@ -1,11 +1,12 @@ from __future__ import annotations import math +import random from functools import partial from typing import TYPE_CHECKING 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 @@ -149,6 +150,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, @@ -243,6 +262,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), } @@ -557,3 +579,61 @@ def filter_dict(d: dict, keep: set | Sequence, *, recursive: bool = False) -> di 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