Convert more nodes to the new input types system

Make WaveletCFG less spammy in verbose mode
WaveletCFG will pass sigmas and other information to latent operations that support it
Add SonarCustomNoiseParameters node
Add replace/replace_keepsign/replace_avoidsign quantile norm modes
This commit is contained in:
blepping
2025-07-29 13:04:03 -06:00
parent 849b7266e7
commit ee6410523e
11 changed files with 732 additions and 665 deletions
+2 -1
View File
@@ -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
+4 -3
View File
@@ -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,
+7 -2
View File
@@ -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(
+91 -186
View File
@@ -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"<FRUXConfig: {meh}>"
class FreeUExtremeNode(metaclass=IntegratedNode):
class FreeUExtremeNode:
DESCRIPTION = "Main FreeU Extreme node. Allows patching a model with the FreeU (V2) effect with more control."
RETURN_TYPES = ("MODEL",)
FUNCTION = "go"
CATEGORY = "model_patches"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": (
"MODEL",
{
"tooltip": "Model to patch.",
},
),
"cpu_fft": (
"BOOLEAN",
{
"default": False,
"tooltip": "Controls whether to perform FFT calculations on the CPU. May be necessary for some GPUs that don't have native support for FFT operations at the cost of performance.",
},
),
},
"optional": {
"input_config": (
"FRUX_CONFIG",
{
"tooltip": "Allows specifying configuration for input blocks.",
},
),
"middle_config": (
"FRUX_CONFIG",
{
"tooltip": "Allows specifying configuration for middle blocks.",
},
),
"output_config": (
"FRUX_CONFIG",
{
"tooltip": "Allows specifying configuration for output blocks.",
},
),
},
}
INPUT_TYPES = (
SonarInputTypes()
.req_model(tooltip="Model to patch.")
.req_bool_cpu_fft(
tooltip="Controls whether to perform FFT calculations on the CPU. May be necessary for some GPUs that don't have native support for FFT )operations at the cost of performance.",
)
.opt_field_input_config(
"FRUX_CONFIG",
tooltip="Allows specifying configuration for input blocks.",
)
.opt_field_middle_config(
"FRUX_CONFIG",
tooltip="Allows specifying configuration for middle blocks.",
)
.opt_field_output_config(
"FRUX_CONFIG",
tooltip="Allows specifying configuration for output blocks.",
)
)
@classmethod
def go(
+83 -170
View File
@@ -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"):
+10 -12
View File
@@ -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."
+140 -14
View File
@@ -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,
}
+102 -178
View File
@@ -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(
+96 -94
View File
@@ -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,
}
+116 -4
View File
@@ -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,
+81 -1
View File
@@ -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