6 Commits
Author SHA1 Message Date
blepping 650467ce97 Add SonarNestedNoise node
Add centering and redistribute mode to quantile normalization features
Various cleanups
2026-08-17 05:15:32 -06:00
blepping bba5bf25e9 Better handling for nested AV latents in the SONAR_CUSTOM_NOISE to NOISE node 2026-08-07 13:14:58 -06:00
blepping 3a753c1a8b How do I hold all these dumb changes? 2026-06-23 11:17:33 -06:00
blepping ec7def5723 Sync changes 2026-03-02 02:50:41 -07:00
blepping 4ec5970128 Phase 2 2025-08-15 17:04:48 -06:00
blepping e4b05c506d Phase 1 2025-08-14 17:36:11 -06:00
29 changed files with 7379 additions and 4270 deletions
+2 -2
View File
@@ -467,8 +467,8 @@ Some modes act as wrappers to other modes. All modes will just ignore parameters
Modes listed with the defaults for parameters they support. These modes also support `dscale` which defaults to 1 and can be used to manually adjust the scale of the mode result.
* `euclidean`
* `manhatten`
* `euclidean` - Default mode, uses Euclidean distances.
* `manhatten` - Uses Manhatten distances.
* `chebyshev`
* `minkowsi:p=3.0`
* `quadratic`
+4 -3
View File
@@ -4,9 +4,10 @@ import contextlib
import importlib
import sys
from functools import partial
from typing import TYPE_CHECKING, Callable, NamedTuple
from typing import TYPE_CHECKING, Any, NamedTuple
if TYPE_CHECKING:
from collections.abc import Callable
from types import ModuleType
@@ -83,7 +84,7 @@ class Integrations:
class SonarIntegrations(Integrations):
def __init__(self, *args: list, **kwargs: dict):
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
self.register_integration("bleh", "ComfyUI-bleh", self.bleh_integration)
self.register_integration(
@@ -114,7 +115,7 @@ MODULES = SonarIntegrations()
class IntegratedNode(type):
@staticmethod
def wrap_INPUT_TYPES(orig_method: Callable, *args: list, **kwargs: dict) -> dict:
def wrap_INPUT_TYPES(orig_method: Callable, *args: Any, **kwargs: Any) -> dict:
MODULES.initialize()
return orig_method(*args, **kwargs)
+19 -13
View File
@@ -2,14 +2,14 @@ from __future__ import annotations
import math
import random
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any
import torch
from . import utils
if TYPE_CHECKING:
from types import Sequence
from collections.abc import Sequence
class SonarLatentOperation:
@@ -34,9 +34,9 @@ class SonarLatentOperation:
def call_op(
self,
t: torch.Tensor,
*args: list,
*args: Any,
op=None,
**kwargs: dict,
**kwargs: Any,
) -> torch.Tensor:
if op is None:
op = self.op
@@ -51,7 +51,7 @@ class SonarLatentOperation:
latent: torch.Tensor,
*,
sigma: torch.Tensor | float | None = None,
**kwargs: dict,
**kwargs: Any,
) -> torch.Tensor:
if not self.enabled(sigma=sigma):
return latent
@@ -64,16 +64,18 @@ class SonarLatentOperationAdvanced(SonarLatentOperation):
*,
blend_mode: str,
blend_strength: float,
blend_strategy: str,
input_multiplier: float,
output_multiplier: float,
difference_multiplier: float,
ops: Sequence,
op_alt=None,
**kwargs: dict,
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
self.blend_function = utils.BLENDING_MODES[blend_mode]
self.blend_strength = blend_strength
self.blend_strategy = blend_strategy
self.input_multiplier = input_multiplier
self.output_multiplier = output_multiplier
self.difference_multiplier = difference_multiplier
@@ -85,7 +87,7 @@ class SonarLatentOperationAdvanced(SonarLatentOperation):
latent: torch.Tensor,
*,
sigma: torch.Tensor | float | None = None,
**kwargs: dict,
**kwargs: Any,
) -> torch.Tensor:
t = latent
enabled = self.enabled(sigma)
@@ -103,19 +105,23 @@ class SonarLatentOperationAdvanced(SonarLatentOperation):
) - t
if self.difference_multiplier != 1.0:
diff *= self.difference_multiplier
return self.blend_function(t, diff, self.blend_strength)
if self.blend_strategy == "difference":
return self.blend_function(t, diff, self.blend_strength)
if self.blend_strategy == "result":
return self.blend_function(t, t + diff, self.blend_strength)
raise ValueError(f"Unknown blend strategy: {self.blend_strategy}")
class SonarLatentOperationNoise(SonarLatentOperation):
def __init__(
self,
*args: list,
*args: Any,
custom_noise,
scale_to_sigma: bool = False,
cpu_noise: bool = False,
normalize: bool = True,
lazy_noise_sampler: bool = False,
**kwargs: dict,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
self.custom_noise = custom_noise
@@ -131,7 +137,7 @@ class SonarLatentOperationNoise(SonarLatentOperation):
latent: torch.Tensor,
*,
sigma: torch.Tensor | float | None = None,
**kwargs: dict,
**kwargs: Any,
) -> torch.Tensor:
t = latent
enabled = self.enabled(sigma)
@@ -187,12 +193,12 @@ class SonarLatentOperationNoise(SonarLatentOperation):
class SonarLatentOperationSetSeed(SonarLatentOperation):
def __init__(self, *args: list, seed: int, restore_rng_state: bool, **kwargs: dict):
def __init__(self, *args: Any, seed: int, restore_rng_state: bool, **kwargs: Any):
super().__init__(*args, **kwargs)
self.seed = seed
self.restore_rng_state = restore_rng_state
def __call__(self, *args: list, **kwargs: dict) -> torch.Tensor:
def __call__(self, *args: Any, **kwargs: Any) -> torch.Tensor:
if self.restore_rng_state:
pyrandst = random.getstate()
torchrandst = torch.random.get_rng_state()
+42 -38
View File
@@ -25,11 +25,11 @@ NOISE_INPUT_TYPES = frozenset(("SONAR_CUSTOM_NOISE", "OCS_NOISE"))
if not HAVE_COMFY_UNION_TYPE:
class Wildcard(str): # noqa: FURB189
class Wildcard(str):
__slots__ = ("whitelist",)
@classmethod
def __new__(cls, s, *args: list, whitelist=None, **kwargs: dict):
def __new__(cls, s, *args: Any, whitelist=None, **kwargs: Any):
result = super().__new__(s, *args, **kwargs)
result.whitelist = whitelist
return result
@@ -48,17 +48,19 @@ NOISE_INPUT_TYPES_HINT = (
class SonarInputCollection(InputCollection):
def __init__(self, *args: list, **kwargs: dict):
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
self._DELEGATE_KEYS = self._DELEGATE_KEYS | frozenset(( # noqa: PLR6104
"customnoise",
"floatpct",
"normalizetristate",
"selectblend",
"selectnoise",
"selectscalemode",
"yaml",
))
self._DELEGATE_KEYS = self._DELEGATE_KEYS | frozenset(
(
"customnoise",
"floatpct",
"normalizetristate",
"selectblend",
"selectnoise",
"selectscalemode",
"yaml",
),
)
def yaml(
self,
@@ -68,7 +70,7 @@ class SonarInputCollection(InputCollection):
placeholder="# YAML or JSON here",
dynamicPrompts=False, # noqa: N803
multiline=True,
**kwargs: dict,
**kwargs: Any,
):
return self.field(
name,
@@ -87,7 +89,7 @@ class SonarInputCollection(InputCollection):
default="lerp",
insert_modes=(),
tooltip="Mode used for blending. If you have ComfyUI-bleh then you will have access to many more blend modes.",
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
if not MODULES.initialized:
raise RuntimeError(
@@ -108,7 +110,7 @@ class SonarInputCollection(InputCollection):
default="nearest-exact",
insert_modes=(),
tooltip="Mode used for scaling. If you have ComfyUI-bleh then you will have access to many more scale modes.",
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
if not MODULES.initialized:
raise RuntimeError(
@@ -129,7 +131,7 @@ class SonarInputCollection(InputCollection):
default="gaussian",
insert_types=(),
tooltip="Sets the type of noise.",
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
return self.field(
name,
@@ -144,7 +146,7 @@ class SonarInputCollection(InputCollection):
name: str,
add_hint: bool = True, # noqa: FBT001
tooltip="Allows connecting a custom noise chain.",
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
if add_hint:
tooltip = f"{tooltip}\n{NOISE_INPUT_TYPES_HINT}"
@@ -156,7 +158,7 @@ class SonarInputCollection(InputCollection):
*,
default="default",
tooltip="Controls whether noise is normalized to 1.0 strength.",
**kwargs: dict,
**kwargs: Any,
):
return self.field(
name,
@@ -166,14 +168,14 @@ class SonarInputCollection(InputCollection):
**kwargs,
)
def floatpct(self, name: str, *, min=0.0, max=1.0, **kwargs: dict): # noqa: A002
def floatpct(self, name: str, *, min=0.0, max=1.0, **kwargs: Any): # noqa: A002
return self.float(name=name, min=min, max=max, **kwargs)
class SonarInputTypes(InputTypes):
_NO_REPLACE = True
def __init__(self, *args: list, **kwargs: dict):
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(
*args,
collection_class=SonarInputCollection,
@@ -184,7 +186,7 @@ class SonarInputTypes(InputTypes):
class SonarLazyInputTypes(LazyInputTypes):
_NO_REPLACE = True
def __init__(self, *args: list, initializers=(MODULES.initialize,), **kwargs: dict):
def __init__(self, *args: Any, initializers=(MODULES.initialize,), **kwargs: Any):
super().__init__(
*args,
initializers=initializers,
@@ -204,20 +206,22 @@ class SonarCustomNoiseNodeBase(metaclass=IntegratedNode):
raise NotImplementedError
INPUT_TYPES = SonarLazyInputTypes(
lambda *, include_rescale=True, include_chain=True: SonarInputTypes()
.req_float_factor(
default=1.0,
tooltip="Scaling factor for the generated noise of this type.",
)
.req_float_rescale(
_skip=not include_rescale,
default=0.0,
min=0.0,
tooltip="When non-zero, this custom noise item and other custom noise items items connected to it will have their factor scaled to add up to the specified rescale value. When set to 0, rescaling is disabled.",
)
.opt_customnoise_sonar_custom_noise_opt(
_skip=not include_chain,
tooltip="Optional input for more custom noise items.",
lambda *, include_rescale=True, include_chain=True: (
SonarInputTypes()
.req_float_factor(
default=1.0,
tooltip="Scaling factor for the generated noise of this type.",
)
.req_float_rescale(
_skip=not include_rescale,
default=0.0,
min=0.0,
tooltip="When non-zero, this custom noise item and other custom noise items items connected to it will have their factor scaled to add up to the specified rescale value. When set to 0, rescaling is disabled.",
)
.opt_customnoise_sonar_custom_noise_opt(
_skip=not include_chain,
tooltip="Optional input for more custom noise items.",
)
),
initializers=(),
)
@@ -227,7 +231,7 @@ class SonarCustomNoiseNodeBase(metaclass=IntegratedNode):
factor=1.0,
rescale=0.0,
sonar_custom_noise_opt=None,
**kwargs: dict[str, Any],
**kwargs: Any[str, Any],
):
nis = (
sonar_custom_noise_opt.clone()
@@ -240,7 +244,7 @@ class SonarCustomNoiseNodeBase(metaclass=IntegratedNode):
class NoiseChainInputTypes(SonarInputTypes):
def __init__(self, *, parent=SonarCustomNoiseNodeBase, **kwargs: dict):
def __init__(self, *, parent=SonarCustomNoiseNodeBase, **kwargs: Any):
super().__init__(parent=parent, **kwargs)
@@ -251,7 +255,7 @@ class NoiseNoChainInputTypes(SonarInputTypes):
parent=SonarCustomNoiseNodeBase,
parent_args=(),
parent_kwargs=None,
**kwargs: dict,
**kwargs: Any,
):
super().__init__(
parent=parent,
+59 -50
View File
@@ -3,29 +3,38 @@ from __future__ import annotations
from copy import deepcopy
from functools import partial
from typing import Callable, TypeVar
from typing import TYPE_CHECKING, Any, TypeVar
if TYPE_CHECKING:
from collections.abc import Callable
bi_int = int
bi_bool = bool
bi_float = float
class InputCollection:
_DELEGATE_KEYS = frozenset((
"bool",
"boolean",
"clip",
"conditioning",
"field",
"float",
"image",
"int",
"latent",
"model",
"sampler",
"seed",
"sigmas",
"string",
"vae",
))
_DELEGATE_KEYS = frozenset(
(
"bool",
"boolean",
"clip",
"conditioning",
"field",
"float",
"image",
"int",
"latent",
"model",
"sampler",
"seed",
"sigmas",
"string",
"vae",
),
)
def __init__(self, **kwargs: dict):
def __init__(self, **kwargs: Any):
self.fields = kwargs
def __getattr__(self, key: str):
@@ -42,10 +51,10 @@ class InputCollection:
def clone(self):
return InputCollection(**self.to_dict())
def __len__(self) -> int:
def __len__(self) -> bi_int:
return len(self.fields)
def __contains__(self, key: str) -> bool:
def __contains__(self, key: str) -> bi_bool:
return key in self.fields
def field(
@@ -53,8 +62,8 @@ class InputCollection:
name: str,
type: str | tuple,
*,
_skip: bool = False,
**kwargs: dict,
_skip: bi_bool = False,
**kwargs: Any,
) -> InputCollection:
if not _skip:
self.fields[name] = (type,) if not kwargs else (type, kwargs)
@@ -63,7 +72,7 @@ class InputCollection:
def string(
self,
name: str,
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
return self.field(name, "STRING", **kwargs)
@@ -71,11 +80,11 @@ class InputCollection:
self,
name: str,
*,
step: float = 0.001,
min: float = -10000.0,
max: float = 10000.0,
round: bool = False,
**kwargs: dict,
step: bi_float = 0.001,
min: bi_float = -10000.0,
max: bi_float = 10000.0,
round: bi_bool = False,
**kwargs: Any,
) -> InputCollection:
return self.field(
name,
@@ -91,9 +100,9 @@ class InputCollection:
self,
name: str,
*,
min: float = -10000,
max: float = 10000,
**kwargs: dict,
min: bi_float = -10000,
max: bi_float = 10000,
**kwargs: Any,
) -> InputCollection:
return self.field(
name,
@@ -106,22 +115,22 @@ class InputCollection:
def bool(
self,
name: str,
default: bool = False,
**kwargs: dict,
default: bi_bool = False,
**kwargs: Any,
) -> InputCollection:
return self.field(name, "BOOLEAN", default=default, **kwargs)
boolean = bool
boolean = bool # noqa: A003
def seed(
self,
name: str = "seed",
*,
default: int = 0,
min: int = 0,
max: int = 0xFFFFFFFFFFFFFFFF,
default: bi_int = 0,
min: bi_int = 0,
max: bi_int = 0xFFFFFFFFFFFFFFFF,
tooltip="Seed to use for generated noise",
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
return self.int(
name,
@@ -132,32 +141,32 @@ class InputCollection:
**kwargs,
)
def image(self, name: str = "image", **kwargs: dict) -> InputCollection:
def image(self, name: str = "image", **kwargs: Any) -> InputCollection:
return self.field(name, "IMAGE", **kwargs)
def latent(self, name: str = "latent", **kwargs: dict) -> InputCollection:
def latent(self, name: str = "latent", **kwargs: Any) -> InputCollection:
return self.field(name, "LATENT", **kwargs)
def conditioning(
self,
name: str = "conditioning",
**kwargs: dict,
**kwargs: Any,
) -> InputCollection:
return self.field(name, "CONDITIONING", **kwargs)
def model(self, name: str = "model", **kwargs: dict) -> InputCollection:
def model(self, name: str = "model", **kwargs: Any) -> InputCollection:
return self.field(name, "MODEL", **kwargs)
def sigmas(self, name: str = "sigmas", **kwargs: dict) -> InputCollection:
def sigmas(self, name: str = "sigmas", **kwargs: Any) -> InputCollection:
return self.field(name, "SIGMAS", **kwargs)
def sampler(self, name: str = "sampler", **kwargs: dict) -> InputCollection:
def sampler(self, name: str = "sampler", **kwargs: Any) -> InputCollection:
return self.field(name, "SAMPLER", **kwargs)
def clip(self, name: str = "clip", **kwargs: dict) -> InputCollection:
def clip(self, name: str = "clip", **kwargs: Any) -> InputCollection:
return self.field(name, "CLIP", **kwargs)
def vae(self, name: str = "vae", **kwargs: dict) -> InputCollection:
def vae(self, name: str = "vae", **kwargs: Any) -> InputCollection:
return self.field(name, "VAE", **kwargs)
@@ -226,7 +235,7 @@ class InputTypes:
errstr = f"Unknown attribute {key} for InputTypes"
raise AttributeError(errstr)
def wrapper(*args: list, **kwargs: dict):
def wrapper(*args: Any, **kwargs: Any):
meth(*args, **kwargs)
return self
@@ -240,7 +249,7 @@ class LazyInputTypes:
self.builder = builder
self.initializers = initializers
def get_input_types(self, *args: list, **kwargs: dict):
def get_input_types(self, *args: Any, **kwargs: Any):
if args or kwargs:
args = tuple(args)
cache_key = (args, tuple(kwargs.items()))
@@ -259,5 +268,5 @@ class LazyInputTypes:
self._input_types_params[cache_key] = result
return result
def __call__(self, *args: list, **kwargs: dict) -> dict:
def __call__(self, *args: Any, **kwargs: Any) -> dict:
return self.get_input_types(*args, **kwargs)()
+204 -165
View File
@@ -27,93 +27,95 @@ class SonarApplyLatentOperationCFG(metaclass=IntegratedNode):
FUNCTION = "go"
INPUT_TYPES = SonarLazyInputTypes(
lambda: SonarInputTypes()
.req_model()
.req_field_mode(
(
"cond_sub_uncond",
"denoised_sub_uncond",
"uncond_sub_cond",
"denoised",
"cond",
"uncond",
"model_input",
),
default="cond_sub_uncond",
tooltip="cond_sub_uncond is what ComfyUI's latent operations use. The non-sub_uncond modes likely won't work with pred_flip mode enabled. If you have anything but the denoised options selected, this will use pre-CFG, otherwise it will use post-CFG (unless you are using model_input).",
)
.req_bool_pred_flip_mode(
tooltip="Lets you try to apply the latent operation to the noise prediction rather than the image prediction. Doesn't work properly with the non-sub_uncond modes. No real reason it should be better, just something you can try. Note: The noise prediction gets scaled by the sigma first, in case that's useful information.",
)
.req_bool_require_uncond(
tooltip="When enabled, the operation will be skipped if uncond is unavailable. This will also happen if you choose a mode that requires uncond.",
)
.req_float_start_sigma(
default=-1.0,
min=-1.0,
tooltip="First sigma the effect becomes active. You can set a negative value here to use whatever the model's maximum sigma is.",
)
.req_float_end_sigma(
default=0.0,
min=0.0,
tooltip="Last sigma the effect is active.",
)
.req_selectblend_blend_mode(
tooltip="Controls how the output of the latent operation is blended with the original result.",
)
.req_float_blend_strength(
default=0.5,
tooltip="Strength of the blend. For a normal blend mode like LERP, 1.0 means use 100% of the output from the latent operation, 0.0 means use none of it and only the original value. Note: Blending is applied to the final result of the operations unless you enable immediate_blend, in other words operation_2 sees a full unblended result from operation_1.",
)
.req_field_blend_scale_mode(
(
"none",
"reverse_sampling",
"sampling",
"reverse_enabled_range",
"enabled_range",
"sampling_sin",
"enabled_range_sin",
),
default="reverse_sampling",
tooltip="Can be used to scale the blend strength over time. Basically works like blend_strength * scale_factor (see below)\nnone: Just uses the blend_strength you have set.\nreverse_sampling: The opposite of the model sampling percent, so if you're making a new generation, the beginning of sampling will be 1.0 and the end will be 0.0. The recommended option as applying these operations usually works better toward the beginning of sampling.\nsampling: Same as reverse_sampling, except the beginning will be 0.0 and the end will be 1.0.\nreverse_enabled_range: Flipped percentage of the range between start_sigma and end_sigma.\nenabled_range: Percentage of the range between start_sigma and end_sigma.\nsampling_sin: Uses the sampling percentage with the sine function such that blend_strength will hit the peak value in the middle of the range.\nenabled_range_sin: Similar to sampling_sin except it applies to the percentage of the enabled range.",
)
.req_float_blend_scale_offset(
default=0.0,
min=-1.0,
max=1.0,
tooltip="Only applies when blend_scale_mode is not none. Adds the offset to the calculated percentage and then clamps it to be between blend_scale_min and blend_scale_max.",
)
.req_float_blend_scale_min(
default=0.0,
tooltip="Only applies when blend_scale_mode is not none. Minimum value for the blend scale percentage. Many blend modes don't tolerate negative values here.",
)
.req_float_blend_scale_max(
default=1.0,
tooltip="Only applies when blend_scale_mode is not none. Maximum value for the blend scale percentage. Many blend modes don't tolerate values over 1.0 here.",
)
.req_bool_immediate_blend(
tooltip="You can enable this to do blending immediately after each latent operation is called. Mainly affects the case where you have multiple latent operations connected.",
)
.opt_field_operation_1(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_2(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_3(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_4(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_5(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
lambda: (
SonarInputTypes()
.req_model()
.req_field_mode(
(
"cond_sub_uncond",
"denoised_sub_uncond",
"uncond_sub_cond",
"denoised",
"cond",
"uncond",
"model_input",
),
default="cond_sub_uncond",
tooltip="cond_sub_uncond is what ComfyUI's latent operations use. The non-sub_uncond modes likely won't work with pred_flip mode enabled. If you have anything but the denoised options selected, this will use pre-CFG, otherwise it will use post-CFG (unless you are using model_input).",
)
.req_bool_pred_flip_mode(
tooltip="Lets you try to apply the latent operation to the noise prediction rather than the image prediction. Doesn't work properly with the non-sub_uncond modes. No real reason it should be better, just something you can try. Note: The noise prediction gets scaled by the sigma first, in case that's useful information.",
)
.req_bool_require_uncond(
tooltip="When enabled, the operation will be skipped if uncond is unavailable. This will also happen if you choose a mode that requires uncond.",
)
.req_float_start_sigma(
default=-1.0,
min=-1.0,
tooltip="First sigma the effect becomes active. You can set a negative value here to use whatever the model's maximum sigma is.",
)
.req_float_end_sigma(
default=0.0,
min=0.0,
tooltip="Last sigma the effect is active.",
)
.req_selectblend_blend_mode(
tooltip="Controls how the output of the latent operation is blended with the original result.",
)
.req_float_blend_strength(
default=0.5,
tooltip="Strength of the blend. For a normal blend mode like LERP, 1.0 means use 100% of the output from the latent operation, 0.0 means use none of it and only the original value. Note: Blending is applied to the final result of the operations unless you enable immediate_blend, in other words operation_2 sees a full unblended result from operation_1.",
)
.req_field_blend_scale_mode(
(
"none",
"reverse_sampling",
"sampling",
"reverse_enabled_range",
"enabled_range",
"sampling_sin",
"enabled_range_sin",
),
default="reverse_sampling",
tooltip="Can be used to scale the blend strength over time. Basically works like blend_strength * scale_factor (see below)\nnone: Just uses the blend_strength you have set.\nreverse_sampling: The opposite of the model sampling percent, so if you're making a new generation, the beginning of sampling will be 1.0 and the end will be 0.0. The recommended option as applying these operations usually works better toward the beginning of sampling.\nsampling: Same as reverse_sampling, except the beginning will be 0.0 and the end will be 1.0.\nreverse_enabled_range: Flipped percentage of the range between start_sigma and end_sigma.\nenabled_range: Percentage of the range between start_sigma and end_sigma.\nsampling_sin: Uses the sampling percentage with the sine function such that blend_strength will hit the peak value in the middle of the range.\nenabled_range_sin: Similar to sampling_sin except it applies to the percentage of the enabled range.",
)
.req_float_blend_scale_offset(
default=0.0,
min=-1.0,
max=1.0,
tooltip="Only applies when blend_scale_mode is not none. Adds the offset to the calculated percentage and then clamps it to be between blend_scale_min and blend_scale_max.",
)
.req_float_blend_scale_min(
default=0.0,
tooltip="Only applies when blend_scale_mode is not none. Minimum value for the blend scale percentage. Many blend modes don't tolerate negative values here.",
)
.req_float_blend_scale_max(
default=1.0,
tooltip="Only applies when blend_scale_mode is not none. Maximum value for the blend scale percentage. Many blend modes don't tolerate values over 1.0 here.",
)
.req_bool_immediate_blend(
tooltip="You can enable this to do blending immediately after each latent operation is called. Mainly affects the case where you have multiple latent operations connected.",
)
.opt_field_operation_1(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_2(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_3(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_4(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_5(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
),
)
@@ -315,7 +317,7 @@ class SonarApplyLatentOperationCFG(metaclass=IntegratedNode):
class SonarLatentOperationQuantileFilter(SonarQuantileFilteredNoiseNode):
DESCRIPTION = "Allows applying a quantile normalization function to the latent during sampling. Can be used with Sonar SonarApplyLatentOperationCFG. The just copies most of the parameters from the other quantile normalization node where it talks to 'noise', this will apply to whatever you're applying the latent operation to (denoised, uncond, etc)."
DESCRIPTION = "Allows applying a quantile normalization function to the latent during sampling. Can be used with Sonar SonarApplyLatentOperationCFG. The just copies most of the parameters from the other quantile normalization node. When it mentions 'noise' it will affect whatever you're applying the latent operation to (denoised, uncond, etc)."
RETURN_TYPES = ("LATENT_OPERATION",)
CATEGORY = "latent/advanced/operations"
@@ -324,7 +326,13 @@ class SonarLatentOperationQuantileFilter(SonarQuantileFilteredNoiseNode):
result = super().INPUT_TYPES()
result.pop("optional", None)
reqparams = result["required"]
for k in ("custom_noise", "normalize", "normalize_noise", "factor"):
for k in (
"custom_noise",
"reference_noise_opt",
"normalize",
"normalize_noise",
"factor",
):
reqparams.pop(k, None)
return result
@@ -338,7 +346,16 @@ class SonarLatentOperationQuantileFilter(SonarQuantileFilteredNoiseNode):
norm_power: float,
norm_factor: float,
strategy: str,
norm_power_in: float,
sign_mode: str,
abs_quantiles: bool,
only_outliers: bool,
manual_quantiles: str,
zero_mean_scale: str,
):
# TODO: Support an optional reference LATENT_OPERATION.
zms, rms, zrs = cls._parse_mean_scales(zero_mean_scale)
nq_lo, nq_hi = cls._parse_manual_quantiles(manual_quantiles)
qnorm_filter = functools.partial(
utils.quantile_normalize,
quantile=quantile,
@@ -347,6 +364,15 @@ class SonarLatentOperationQuantileFilter(SonarQuantileFilteredNoiseNode):
nq_fac=norm_factor,
pow_fac=norm_power,
strategy=strategy,
pow_fac_in=norm_power_in,
sign_mode=sign_mode,
abs_quantiles=abs_quantiles,
only_outliers=only_outliers,
nq_lo=nq_lo,
nq_hi=nq_hi,
zero_mean_scale=zms,
restore_mean_scale=rms,
zero_result_mean_scale=zrs,
)
return (SonarLatentOperation(op=lambda latent: qnorm_filter(latent)),) # noqa: PLW0108
@@ -360,60 +386,67 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
FUNCTION = "go"
INPUT_TYPES = SonarLazyInputTypes(
lambda: SonarInputTypes()
.req_field_operation(
"LATENT_OPERATION",
tooltip="Latent operation to apply.",
)
.req_float_start_sigma(
default=-1.0,
min=-1.0,
tooltip="First sigma the effect becomes active. You can set a negative value here to use whatever the model's maximum sigma is.",
)
.req_float_end_sigma(
default=0.0,
min=0.0,
tooltip="Last sigma the effect is active.",
)
.req_float_input_multiplier(
default=1.0,
tooltip="Flat multiplier on the input to the latent operation. The multiplied input is *not* used when calculating the difference, it is only passed to the operation.",
)
.req_float_output_multiplier(
default=1.0,
tooltip="Flat multiplier on the output from the latent operation. Occurs before blending or calculating the difference.",
)
.req_float_difference_multiplier(
default=1.0,
tooltip="Flat multiplier on the difference or change from the original that the operation performed. Occurs after output_multiplier and before blending applies.",
)
.req_selectblend_blend_mode(
default="inject",
tooltip="Controls how the change from the operation is combined with the input. The default of inject just adds it scaled by the blend strength. With 1.0 blend strength, this is just using the output from the operation with no change.",
)
.req_float_blend_strength(
default=0.5,
tooltip="Strength of the blend.",
)
.opt_field_operation_alt(
"LATENT_OPERATION",
tooltip="Optional alternative operation that will be used when the primary one isn't enabled. May be useful in a case when you want one operation between sigma 1.0 and 0.5 and then a difference operation for lower sigmas which is kind of annoying to specify manually (you'd need to do something like configure another operation to start at 0.499999 or something).",
)
.opt_field_operation_2(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_3(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_4(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_5(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
lambda: (
SonarInputTypes()
.req_float_start_sigma(
default=-1.0,
min=-1.0,
tooltip="First sigma the effect becomes active. You can set a negative value here to use whatever the model's maximum sigma is.",
)
.req_float_end_sigma(
default=0.0,
min=0.0,
tooltip="Last sigma the effect is active.",
)
.req_float_input_multiplier(
default=1.0,
tooltip="Flat multiplier on the input to the latent operation. The multiplied input is *not* used when calculating the difference, it is only passed to the operation.",
)
.req_float_output_multiplier(
default=1.0,
tooltip="Flat multiplier on the output from the latent operation. Occurs before blending or calculating the difference.",
)
.req_float_difference_multiplier(
default=1.0,
tooltip="Flat multiplier on the difference or change from the original that the operation performed. Occurs after output_multiplier and before blending applies.",
)
.req_selectblend_blend_mode(
default="inject",
tooltip="Controls how the change from the operation is combined with the input. The default of inject just adds it scaled by the blend strength. With 1.0 blend strength, this is just using the output from the operation with no change.",
)
.req_float_blend_strength(
default=0.5,
tooltip="Strength of the blend.",
)
.req_field_blend_strategy(
("difference", "result"),
default="difference",
tooltip="Controls whether blending occurs with the difference or changed result after the latent operation.",
)
.opt_field_operation(
"LATENT_OPERATION",
tooltip="Latent operation to apply.",
)
.opt_field_operation_alt(
"LATENT_OPERATION",
tooltip="Optional alternative operation that will be used when the primary one isn't enabled. May be useful in a case when you want one operation between sigma 1.0 and 0.5 and then a difference operation for lower sigmas which is kind of annoying to specify manually (you'd need to do something like configure another operation to start at 0.499999 or something).",
)
.opt_field_operation_2(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_3(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_4(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
.opt_field_operation_5(
"LATENT_OPERATION",
tooltip="Optional LATENT_OPERATION. The operations will be applied in sequence.",
)
),
)
@@ -421,7 +454,6 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
def go(
cls,
*,
operation,
start_sigma: float,
end_sigma: float,
input_multiplier: float,
@@ -429,6 +461,8 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
difference_multiplier: float,
blend_mode: str,
blend_strength: float,
blend_strategy: str,
operation=None,
operation_alt=None,
operation_2=None,
operation_3=None,
@@ -456,6 +490,7 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
difference_multiplier=difference_multiplier,
blend_mode=blend_mode,
blend_strength=blend_strength,
blend_strategy=blend_strategy,
),
)
@@ -468,19 +503,21 @@ class SonarLatentOperationNoiseNode(metaclass=IntegratedNode):
FUNCTION = "go"
INPUT_TYPES = SonarLazyInputTypes(
lambda: SonarInputTypes()
.req_customnoise_custom_noise()
.req_bool_scale_to_sigma(tooltip="Scales the noise to the current sigma.")
.req_bool_cpu_noise(
tooltip="Controls whether noise is generated on the CPU or GPU. GPU is usually faster but may change seeds for different models of GPU.",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether the generated noise is normalized.",
)
.req_bool_lazy_noise_sampler(
default=True,
tooltip="When enabled, the latent operation will attempt to cache the noise sampler between calls and only recreate it when necessary. However, there isn't a 100% reliable way for a latent operation to know when sampling starts/ends so if we get it wrong this will lead to non-deterministic generations. I believe the heuristic I'm using to detect this should be reliable but you can disable it if you notice weird results.",
lambda: (
SonarInputTypes()
.req_customnoise_custom_noise()
.req_bool_scale_to_sigma(tooltip="Scales the noise to the current sigma.")
.req_bool_cpu_noise(
tooltip="Controls whether noise is generated on the CPU or GPU. GPU is usually faster but may change seeds for different models of GPU.",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether the generated noise is normalized.",
)
.req_bool_lazy_noise_sampler(
default=True,
tooltip="When enabled, the latent operation will attempt to cache the noise sampler between calls and only recreate it when necessary. However, there isn't a 100% reliable way for a latent operation to know when sampling starts/ends so if we get it wrong this will lead to non-deterministic generations. I believe the heuristic I'm using to detect this should be reliable but you can disable it if you notice weird results.",
)
),
)
@@ -513,14 +550,16 @@ class SonarLatentOperationSetSeedNode(metaclass=IntegratedNode):
FUNCTION = "go"
INPUT_TYPES = SonarLazyInputTypes(
lambda: SonarInputTypes()
.req_field_operation("LATENT_OPERATION")
.req_seed(
tooltip="Seed to set. Note that this is called _every time_ before the operation.",
)
.req_bool_restore_rng_state(
default=False,
tooltip="When enabled, the current RNG state is saved just before calling the operation and restored afterwards. In other words, only the latent operation will see the seed you set. Note: This only handles the PyTorch and Python random module states.",
lambda: (
SonarInputTypes()
.req_field_operation("LATENT_OPERATION")
.req_seed(
tooltip="Seed to set. Note that this is called _every time_ before the operation.",
)
.req_bool_restore_rng_state(
default=False,
tooltip="When enabled, the current RNG state is saved just before calling the operation and restored afterwards. In other words, only the latent operation will see the seed you set. Note: This only handles the PyTorch and Python random module states.",
)
),
)
+338 -249
View File
@@ -4,12 +4,13 @@ import functools
import inspect
import math
import random
from typing import Any, Callable
from typing import TYPE_CHECKING, Any
import numpy as np
import torch
import yaml
from comfy import model_management, samplers
from comfy import utils as comfy_utils
from tqdm import tqdm
from .. import noise, utils
@@ -24,6 +25,14 @@ from .base import (
SonarNormalizeNoiseNodeMixin,
)
if TYPE_CHECKING:
from collections.abc import Callable
try:
from comfy import nested_tensor
except (ModuleNotFoundError, ImportError):
nested_tensor = None
class NoisyLatentLikeNode(metaclass=IntegratedNode):
DESCRIPTION = "Allows generating noise (and optionally adding it) based on a reference latent. Note: For img2img workflows, you will generally want to enable add_to_latent as well as connecting the model and sigmas inputs."
@@ -34,38 +43,40 @@ class NoisyLatentLikeNode(metaclass=IntegratedNode):
FUNCTION = "go"
INPUT_TYPES = SonarLazyInputTypes(
lambda: SonarInputTypes()
.req_selectnoise_noise_type(
tooltip="Sets the type of noise to generate. Has no effect when the custom_noise_opt input is connected.",
)
.req_seed()
.req_latent(tooltip="Latent used as a reference for generating noise.")
.req_float_multiplier(
default=1.0,
tooltip="Multiplier for the strength of the generated noise. Performed after mul_by_sigmas_opt.",
)
.req_bool_add_to_latent(
tooltip="Add the generated noise to the reference latent rather than adding it to an empty latent. Generally should be enabled for img2img workflows.",
)
.req_int_repeat_batch(
default=1,
min=1,
tooltip="Repeats the noise generation the specified number of times. For example, if set to two and your reference latent is also batch two you will get a batch of four as output.",
)
.req_bool_cpu_noise(
default=True,
tooltip="Controls whether noise will be generated on GPU or CPU. Only affects noise types that support GPU generation (maybe only Brownian).",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether the generated noise is normalized to 1.0 strength before scaling. Generally should be left enabled.",
)
.opt_customnoise_custom_noise_opt()
.opt_sigmas_mul_by_sigmas_opt(
tooltip="When connected, will scale the generated noise by the first sigma. Must also connect model_opt to enable.",
)
.opt_model_model_opt(
tooltip="Used when mul_by_sigmas_opt is connected, no effect otherwise.",
lambda: (
SonarInputTypes()
.req_selectnoise_noise_type(
tooltip="Sets the type of noise to generate. Has no effect when the custom_noise_opt input is connected.",
)
.req_seed()
.req_latent(tooltip="Latent used as a reference for generating noise.")
.req_float_multiplier(
default=1.0,
tooltip="Multiplier for the strength of the generated noise. Performed after mul_by_sigmas_opt.",
)
.req_bool_add_to_latent(
tooltip="Add the generated noise to the reference latent rather than adding it to an empty latent. Generally should be enabled for img2img workflows.",
)
.req_int_repeat_batch(
default=1,
min=1,
tooltip="Repeats the noise generation the specified number of times. For example, if set to two and your reference latent is also batch two you will get a batch of four as output.",
)
.req_bool_cpu_noise(
default=True,
tooltip="Controls whether noise will be generated on GPU or CPU. Only affects noise types that support GPU generation (maybe only Brownian).",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether the generated noise is normalized to 1.0 strength before scaling. Generally should be left enabled.",
)
.opt_customnoise_custom_noise_opt()
.opt_sigmas_mul_by_sigmas_opt(
tooltip="When connected, will scale the generated noise by the first sigma. Must also connect model_opt to enable.",
)
.opt_model_model_opt(
tooltip="Used when mul_by_sigmas_opt is connected, no effect otherwise.",
)
),
)
@@ -125,7 +136,7 @@ class NoisyLatentLikeNode(metaclass=IntegratedNode):
sigma_max=sigma_max,
seed=seed,
cpu=cpu_noise,
normalized=normalize,
normalized=False,
)
else:
ns = noise.get_noise_sampler(
@@ -135,7 +146,7 @@ class NoisyLatentLikeNode(metaclass=IntegratedNode):
sigma_max,
seed=seed,
cpu=cpu_noise,
normalized=normalize,
normalized=False,
)
randst = torch.random.get_rng_state()
try:
@@ -146,7 +157,7 @@ class NoisyLatentLikeNode(metaclass=IntegratedNode):
)
finally:
torch.random.set_rng_state(randst)
result = utils.scale_noise(result, multiplier, normalized=True)
result = utils.scale_noise(result, multiplier, normalized=normalize)
if add_to_latent:
result += latent_samples.repeat(
*(repeat_batch if i == 0 else 1 for i in range(latent_samples.ndim)),
@@ -163,80 +174,82 @@ class SonarNoiseImageNode(metaclass=IntegratedNode):
FUNCTION = "go"
INPUT_TYPES = SonarLazyInputTypes(
lambda: SonarInputTypes()
.req_selectnoise_noise_type(
tooltip="Sets the type of noise to generate. Has no effect when the custom_noise_opt input is connected.",
)
.req_seed()
.req_image(tooltip="Image noise will be added to.")
.req_float_noise_min(
default=0.0,
tooltip="Generated noise will be normalized to have values between noise_min and noise_max. If you set them both to the same value then this disables normalization.",
)
.req_float_noise_max(
default=1.0,
tooltip="Generated noise will be normalized to have values between noise_min and noise_max. If you set them both to the same value then this disables normalization.",
)
.req_float_noise_multiplier(
default=0.5,
tooltip="Multiplier for the strength of the generated noise. This is performed after noise_min/max scaling.",
)
.req_field_channel_mode(
(
"RGB",
"RGBA",
"R",
"G",
"B",
"A",
"RA",
"GA",
"BA",
"RG",
"RB",
"GB",
"RGA",
"RBA",
"GBA",
),
default="RGB",
tooltip="RGBA will also add noise to the alpha channel as well if it exists. Only used for 3 or 4 channel images, for other numbers of channels (i.e. one channel) then all channels will be targeted.",
)
.req_selectblend(
insert_modes=("simple_add",),
default="simple_add",
tooltip="Controls how the generated noise is combined with the image. simple_add just adds it and blend_strength is ignored in that case.",
)
.req_float_blend_strength(
default=0.5,
tooltip="Multiplier for the strength of the generated noise.",
)
.req_field_overflow_mode(
("clamp", "rescale"),
default="clamp",
tooltip="When set to clamp, values above/below 0, 1 will be set to those values. When set to rescale, the image values will be rescaled such that the minimum value is 0 and the maximum is 1.",
)
.req_bool_greyscale_mode(
tooltip="When set to clamp, values above/below 0, 1 will be set to those values. When set to rescale, the image values will be rescaled such that the minimum value is 0 and the maximum is 1.",
)
.req_bool_pure_noise_mode(
tooltip="When enabled, the original image is only used for its shape and you will be adding noise to an image full of zeros (black), suitable for creating pure noise images.",
)
.req_field_dtype(
("default", "float32", "float64", "float16", "bfloat16"),
default="default",
tooltip="When set to default it will use the same type as the input tensor (probably float32). You can manually set the dtype if you want, though it likely isn't going to matter. Using dtypes with limited range (float16, bfloat16) isn't recommended.",
)
.req_bool_cpu_noise(
default=True,
tooltip="Controls whether noise will be generated on GPU or CPU.",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether the generated noise is normalized to 1.0 strength before scaling. Generally should be left enabled.",
)
.opt_customnoise_custom_noise_opt(
tooltip="Allows connecting a custom noise chain. When connected, noise_type has no effect.",
lambda: (
SonarInputTypes()
.req_selectnoise_noise_type(
tooltip="Sets the type of noise to generate. Has no effect when the custom_noise_opt input is connected.",
)
.req_seed()
.req_image(tooltip="Image noise will be added to.")
.req_float_noise_min(
default=0.0,
tooltip="Generated noise will be normalized to have values between noise_min and noise_max. If you set them both to the same value then this disables normalization.",
)
.req_float_noise_max(
default=1.0,
tooltip="Generated noise will be normalized to have values between noise_min and noise_max. If you set them both to the same value then this disables normalization.",
)
.req_float_noise_multiplier(
default=0.5,
tooltip="Multiplier for the strength of the generated noise. This is performed after noise_min/max scaling.",
)
.req_field_channel_mode(
(
"RGB",
"RGBA",
"R",
"G",
"B",
"A",
"RA",
"GA",
"BA",
"RG",
"RB",
"GB",
"RGA",
"RBA",
"GBA",
),
default="RGB",
tooltip="RGBA will also add noise to the alpha channel as well if it exists. Only used for 3 or 4 channel images, for other numbers of channels (i.e. one channel) then all channels will be targeted.",
)
.req_selectblend(
insert_modes=("simple_add",),
default="simple_add",
tooltip="Controls how the generated noise is combined with the image. simple_add just adds it and blend_strength is ignored in that case.",
)
.req_float_blend_strength(
default=0.5,
tooltip="Multiplier for the strength of the generated noise.",
)
.req_field_overflow_mode(
("clamp", "rescale"),
default="clamp",
tooltip="When set to clamp, values above/below 0, 1 will be set to those values. When set to rescale, the image values will be rescaled such that the minimum value is 0 and the maximum is 1.",
)
.req_bool_greyscale_mode(
tooltip="When set to clamp, values above/below 0, 1 will be set to those values. When set to rescale, the image values will be rescaled such that the minimum value is 0 and the maximum is 1.",
)
.req_bool_pure_noise_mode(
tooltip="When enabled, the original image is only used for its shape and you will be adding noise to an image full of zeros (black), suitable for creating pure noise images.",
)
.req_field_dtype(
("default", "float32", "float64", "float16", "bfloat16"),
default="default",
tooltip="When set to default it will use the same type as the input tensor (probably float32). You can manually set the dtype if you want, though it likely isn't going to matter. Using dtypes with limited range (float16, bfloat16) isn't recommended.",
)
.req_bool_cpu_noise(
default=True,
tooltip="Controls whether noise will be generated on GPU or CPU.",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether the generated noise is normalized to 1.0 strength before scaling. Generally should be left enabled.",
)
.opt_customnoise_custom_noise_opt(
tooltip="Allows connecting a custom noise chain. When connected, noise_type has no effect.",
)
),
)
@@ -373,83 +386,153 @@ class CustomNOISE:
self.normalize = normalize
self.multiplier = multiplier
def _sample_noise(self, latent_image, seed):
result = self.custom_noise.make_noise_sampler(
latent_image,
None,
None,
seed=seed,
cpu=self.cpu_noise,
normalized=self.normalize,
)(None, None).to(
device="cpu",
dtype=latent_image.dtype,
def _sample_noise(self, latent_image, seed, sampler_idx: int = 0):
if self.multiplier == 0.0:
return torch.zeros_like(latent_image)
n_samplers = len(self.custom_noise)
sampler_idx = sampler_idx % n_samplers
result = (
self.custom_noise[sampler_idx]
.make_noise_sampler(
latent_image,
None,
None,
seed=seed,
cpu=self.cpu_noise,
normalized=self.normalize,
)(None, None)
.to(
device="cpu",
dtype=latent_image.dtype,
)
)
if result.layout != latent_image.layout:
if latent_image.layout == torch.sparse_coo:
return result.to_sparse()
errstr = f"Cannot handle latent layout {type(latent_image.layout).__name__}"
raise NotImplementedError(errstr)
return result if self.multiplier == 1.0 else result.mul_(self.multiplier)
if self.multiplier != 1.0:
result *= self.multiplier
return result
def generate_noise(self, input_latent):
latent_image = input_latent["samples"]
orig_type = type(latent_image)
# print(f"\nNEST? {latent_image.is_nested}, have={nested_tensor is not None}")
if nested_tensor is not None and latent_image.is_nested:
nested = True
nested_parts = latent_image.unbind()
# print(f"NEST: {tuple(p.shape for p in nested_parts)}")
else:
nested = False
nested_parts = (latent_image,)
batch_inds = input_latent.get("batch_index")
torch.manual_seed(self.seed)
random.seed(self.seed)
if self.multiplier == 0.0:
return torch.zeros(
latent_image.shape,
dtype=latent_image.dtype,
layout=latent_image.layout,
device="cpu",
)
if batch_inds is None:
return self._sample_noise(latent_image, self.seed)
unique_inds, inverse_inds = np.unique(batch_inds, return_inverse=True)
result = []
batch_size = latent_image.shape[0]
for idx in range(unique_inds[-1] + 1):
noise = self._sample_noise(
latent_image[idx % batch_size].unsqueeze(0),
self.seed + idx,
noise_parts = tuple(
self._sample_noise(nested_parts[i], self.seed, sampler_idx=i)
for i in range(len(nested_parts))
)
if idx in unique_inds:
result.append(noise)
return torch.cat(tuple(result[i] for i in inverse_inds), axis=0)
return (
orig_type(comfy_utils.pack_latents(noise_parts)[0])
if nested
else noise_parts[0]
)
batch_size = latent_image.shape[0]
unique_inds, inverse_inds = np.unique(batch_inds, return_inverse=True)
use_idxs = (idx for idx in range(unique_inds[-1] + 1) if idx in unique_inds)
use_idxs = {idx: inverse_inds[uidx] for uidx, idx in enumerate(use_idxs)}
n_use_idxs = len(use_idxs)
result_parts = tuple(
torch.empty(
(n_use_idxs, *np.shape[1:]),
dtype=latent_image.dtype,
device=latent_image.device,
)
for np in nested_parts
)
for idx in range(unique_inds[-1] + 1):
sample_idx = idx % batch_size
for nidx in range(len(nested_parts)):
sample = nested_parts[nidx][sample_idx].unsqueeze(0)
noise = self._sample_noise(sample, self.seed + idx, sampler_idx=nidx)
batch_out_idx = use_idxs.get(idx)
if batch_out_idx is not None:
result = result_parts[nidx]
result[batch_out_idx : batch_out_idx + 1] = noise[:1]
return (
orig_type(comfy_utils.pack_latents(result_parts)[0])
if nested
else result_parts[0]
)
class SonarToComfyNOISENode(metaclass=IntegratedNode):
DESCRIPTION = "Allows converting SONAR_CUSTOM_NOISE to NOISE (used by SamplerCustomAdvanced and possibly other custom samplers). NOTE: Does not work with noise types that depend on sigma (Brownian, ScheduledNoise, etc)."
DESCRIPTION = "Allows converting SONAR_CUSTOM_NOISE to NOISE (used by SamplerCustomAdvanced and possibly other custom samplers). The extra alt inputs are used if the latent is a nested tensor and ignored otherwise. Audio/video models like LTX and MiniMax H3 used nested tensors (order video then audio). Connected custom noise inputs will be used in order. NOTE: This node does not work with noise types that depend on sigma (Brownian, ScheduledNoise, etc) unless you manually set a sigma via other nodes."
RETURN_TYPES = ("NOISE",)
CATEGORY = "sampling/custom_sampling/noise"
FUNCTION = "go"
INPUT_TYPES = SonarLazyInputTypes(
lambda: SonarInputTypes()
.req_customnoise_custom_noise(
tooltip="Custom noise type to convert.",
)
.req_seed(tooltip="Seed to use for generated noise.")
.req_bool_cpu_noise(
default=True,
tooltip="Controls whether noise is generated on CPU or GPU.",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether generated noise is normalized to 1.0 strength.",
)
.req_float_multiplier(
default=1.0,
tooltip="Simple multiplier applied to noise after all other scaling and normalization effects. If set to 0, no noise will be generated (same as disabling noise).",
lambda: (
SonarInputTypes()
.req_customnoise_custom_noise(
tooltip="Custom noise type to convert.",
)
.req_seed(tooltip="Seed to use for generated noise.")
.req_bool_cpu_noise(
default=False,
tooltip="Controls whether noise is generated on CPU or GPU.",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether generated noise is normalized to 1.0 strength.",
)
.req_float_multiplier(
default=1.0,
tooltip="Simple multiplier applied to noise after all other scaling and normalization effects. If set to 0, no noise will be generated (same as disabling noise).",
)
.opt_customnoise_alt_custom_noise_1(
tooltip="Optional custom noise. See the node description.",
)
.opt_customnoise_alt_custom_noise_2(
tooltip="Optional custom noise. See the node description.",
)
.opt_customnoise_alt_custom_noise_3(
tooltip="Optional custom noise. See the node description.",
)
),
)
@classmethod
def go(cls, *, custom_noise, seed, cpu_noise=True, normalize=True, multiplier=1.0):
def go(
cls,
*,
custom_noise,
seed,
cpu_noise=True,
normalize=True,
multiplier=1.0,
alt_custom_noise_1=None,
alt_custom_noise_2=None,
alt_custom_noise_3=None,
):
noises = tuple(
cn.clone()
for cn in (
custom_noise,
alt_custom_noise_1,
alt_custom_noise_2,
alt_custom_noise_3,
)
if cn is not None
)
return (
CustomNOISE(
custom_noise,
noises,
seed,
cpu_noise=cpu_noise,
normalize=normalize,
@@ -462,45 +545,47 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode):
DESCRIPTION = "Allows overriding paramaters for a SAMPLER. Only parameters that particular sampler supports will be applied, so for example setting ETA will have no effect for non-ancestral Euler."
INPUT_TYPES = SonarLazyInputTypes(
lambda: SonarInputTypes()
.req_sampler()
.req_float_eta(
default=1.0,
tooltip="Basically controls the ancestralness of the sampler. When set to 0, you will get a non-ancestral (or SDE) sampler.",
)
.req_float_s_noise(
default=1.0,
tooltip="Multiplier for noise added during ancestral or SDE sampling.",
)
.req_float_s_churn(
default=0.0,
tooltip="Churn was the predececessor of ETA. Only used by a few types of samplers (notably Euler non-ancestral). Not used by any ancestral or SDE samplers.",
)
.req_float_r(
default=0.5,
tooltip="Used by dpmpp_sde (and perhaps a few other SDE samplers).",
)
.req_field_sde_solver(
("midpoint", "heun"),
tooltip="Solver used by dpmpp_2m_sde.",
)
.req_bool_cpu_noise(
default=True,
tooltip="Controls whether noise is generated on CPU or GPU.",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether generated noise is normalized to 1.0 strength.",
)
.opt_selectnoise_noise_type(
insert_types=("DEFAULT",),
default="DEFAULT",
tooltip="Noise type used during ancestral or SDE sampling. DEFAULT will use the default for the attached sampler. Only used when the custom noise input is not connected.",
)
.opt_customnoise_custom_noise_opt(
tooltip="Optional input for custom noise used during ancestral or SDE sampling. When connected, the built-in noise_type selector is ignored.",
)
.opt_yaml(),
lambda: (
SonarInputTypes()
.req_sampler()
.req_float_eta(
default=1.0,
tooltip="Basically controls the ancestralness of the sampler. When set to 0, you will get a non-ancestral (or SDE) sampler.",
)
.req_float_s_noise(
default=1.0,
tooltip="Multiplier for noise added during ancestral or SDE sampling.",
)
.req_float_s_churn(
default=0.0,
tooltip="Churn was the predececessor of ETA. Only used by a few types of samplers (notably Euler non-ancestral). Not used by any ancestral or SDE samplers.",
)
.req_float_r(
default=0.5,
tooltip="Used by dpmpp_sde (and perhaps a few other SDE samplers).",
)
.req_field_sde_solver(
("midpoint", "heun"),
tooltip="Solver used by dpmpp_2m_sde.",
)
.req_bool_cpu_noise(
default=True,
tooltip="Controls whether noise is generated on CPU or GPU.",
)
.req_bool_normalize(
default=True,
tooltip="Controls whether generated noise is normalized to 1.0 strength.",
)
.opt_selectnoise_noise_type(
insert_types=("DEFAULT",),
default="DEFAULT",
tooltip="Noise type used during ancestral or SDE sampling. DEFAULT will use the default for the attached sampler. Only used when the custom noise input is not connected.",
)
.opt_customnoise_custom_noise_opt(
tooltip="Optional input for custom noise used during ancestral or SDE sampling. When connected, the built-in noise_type selector is ignored.",
)
.opt_yaml()
),
)
RETURN_TYPES = ("SAMPLER",)
@@ -569,11 +654,11 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode):
model,
x,
sigmas,
*args: list[Any],
*args: Any,
override_sampler_cfg: dict[str, Any] | None = None,
noise_sampler: Callable | None = None,
extra_args: dict[str, Any] | None = None,
**kwargs: dict[str, Any],
**kwargs: Any,
) -> torch.Tensor:
if not override_sampler_cfg:
raise ValueError("Override sampler config missing!")
@@ -629,11 +714,13 @@ class SonarSplitNoiseChainNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNode
DESCRIPTION = "Custom noise type that allows splitting off a new chain. This can be useful if you want a link in the chain to be a blended type."
INPUT_TYPES = SonarLazyInputTypes(
lambda: NoiseChainInputTypes()
.req_normalizetristate_normalize(
tooltip="Controls whether the generated noise is normalized to 1.0 strength.",
)
.opt_customnoise_custom_noise(),
lambda: (
NoiseChainInputTypes()
.req_normalizetristate_normalize(
tooltip="Controls whether the generated noise is normalized to 1.0 strength.",
)
.opt_customnoise_custom_noise()
),
)
@classmethod
@@ -796,50 +883,52 @@ verbose: false
"""
INPUT_TYPES = SonarLazyInputTypes(
lambda _yaml_placeholder=_yaml_placeholder: SonarInputTypes()
.req_model()
.req_float_start_sigma(
default=-1.0,
min=-1.0,
tooltip="First sigma wavelet CFG will be used.",
)
.req_float_end_sigma(
default=0.0,
min=0.0,
tooltip="Last sigma wavelet CFG will be used.",
)
.req_field_fallback_mode(
("existing", "own"),
default="existing",
tooltip="Existing mode uses whatever CFG function existed set when this model patch was applied. Own mode does the CFG calculation on its own. The scale will be whatever you set in your guider or sampler.",
)
.req_selectblend_blend_mode(
tooltip="Controls how the result from wavelet CFG is blended with normal CFG. The default of LERP with strength 1.0 uses 100% wavelet CFG.",
)
.req_float_blend_strength(
default=1.0,
tooltip="Controls how the result from wavelet CFG is blended with normal CFG. The default of LERP with strength 1.0 uses 100% wavelet CFG.",
)
.req_yaml(default=_yaml_placeholder)
.opt_field_operation_cond(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to cond. Note: Latent operations only apply if a rule matches.",
)
.opt_field_operation_uncond(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to uncond. Note: Latent operations only apply if a rule matches.",
)
.opt_field_operation_fallback_cfg(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to the fallback (non-wavelet) CFG result. Note: Latent operations only apply if a rule matches.",
)
.opt_field_operation_wavelet_cfg(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to wavelet CFG result. Note: Latent operations only apply if a rule matches.",
)
.opt_field_operation_result(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to the final result, after wavelet and normal CFG are potentially blended. Note: Latent operations only apply if a rule matches.",
lambda _yaml_placeholder=_yaml_placeholder: (
SonarInputTypes()
.req_model()
.req_float_start_sigma(
default=-1.0,
min=-1.0,
tooltip="First sigma wavelet CFG will be used.",
)
.req_float_end_sigma(
default=0.0,
min=0.0,
tooltip="Last sigma wavelet CFG will be used.",
)
.req_field_fallback_mode(
("existing", "own"),
default="existing",
tooltip="Existing mode uses whatever CFG function existed set when this model patch was applied. Own mode does the CFG calculation on its own. The scale will be whatever you set in your guider or sampler.",
)
.req_selectblend_blend_mode(
tooltip="Controls how the result from wavelet CFG is blended with normal CFG. The default of LERP with strength 1.0 uses 100% wavelet CFG.",
)
.req_float_blend_strength(
default=1.0,
tooltip="Controls how the result from wavelet CFG is blended with normal CFG. The default of LERP with strength 1.0 uses 100% wavelet CFG.",
)
.req_yaml(default=_yaml_placeholder)
.opt_field_operation_cond(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to cond. Note: Latent operations only apply if a rule matches.",
)
.opt_field_operation_uncond(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to uncond. Note: Latent operations only apply if a rule matches.",
)
.opt_field_operation_fallback_cfg(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to the fallback (non-wavelet) CFG result. Note: Latent operations only apply if a rule matches.",
)
.opt_field_operation_wavelet_cfg(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to wavelet CFG result. Note: Latent operations only apply if a rule matches.",
)
.opt_field_operation_result(
"LATENT_OPERATION",
tooltip="Optional latent operation that will be applied to the final result, after wavelet and normal CFG are potentially blended. Note: Latent operations only apply if a rule matches.",
)
),
)
+949 -676
View File
File diff suppressed because it is too large Load Diff
+583 -312
View File
@@ -18,30 +18,32 @@ class SonarAdvancedPyramidNoiseNode(SonarCustomNoiseNodeBase):
)
INPUT_TYPES = SonarLazyInputTypes(
lambda: NoiseChainInputTypes()
.req_field_variant(
(
"highres_pyramid",
"pyramid",
"pyramid_old",
),
default="highres_pyramid",
tooltip="Sets the Pyramid noise variant to generate.",
)
.req_int_iterations(
default=-1,
min=-1,
max=8,
tooltip="When set to -1 will use the variant default.",
)
.req_float_discount(
default=0.0,
tooltip="When set to 0 will use the variant default.",
)
.req_selectscalemode_upscale_mode(
insert_modes=("default",),
default="default",
tooltip="Allows setting the scaling mode for Pyramid noise. Leave on default to use the variant default.",
lambda: (
NoiseChainInputTypes()
.req_field_variant(
(
"highres_pyramid",
"pyramid",
"pyramid_old",
),
default="highres_pyramid",
tooltip="Sets the Pyramid noise variant to generate.",
)
.req_int_iterations(
default=-1,
min=-1,
max=8,
tooltip="When set to -1 will use the variant default.",
)
.req_float_discount(
default=0.0,
tooltip="When set to 0 will use the variant default.",
)
.req_selectscalemode_upscale_mode(
insert_modes=("default",),
default="default",
tooltip="Allows setting the scaling mode for Pyramid noise. Leave on default to use the variant default.",
)
),
)
@@ -75,26 +77,28 @@ class SonarAdvanced1fNoiseNode(SonarCustomNoiseNodeBase):
DESCRIPTION = "Custom noise type that allows specifying parameters for 1f (pink, green, etc) variants."
INPUT_TYPES = SonarLazyInputTypes(
lambda: NoiseChainInputTypes()
.req_float_alpha(
default=0.25,
tooltip="Similar to the advanced power noise node, positive values increase low frequencies (with colorful effects), negative values increase high frequencies.",
)
.req_float_k(
default=1.0,
tooltip="Currently no description of exactly what it does, it's just another knob you can try turning for a different effect.",
)
.req_float_vertical_factor(
default=1.0,
tooltip="Vertical frequency scaling factor.",
)
.req_float_horizontal_factor(
default=1.0,
tooltip="Horizontal frequency scaling factor.",
)
.req_bool_use_sqrt(
default=True,
tooltip="Controls whether to sqrt when dividing the FFT. Negative hfac/wfac won't work when enabled. Turning it off seems to make the parameters have a much stronger effect.",
lambda: (
NoiseChainInputTypes()
.req_float_alpha(
default=0.25,
tooltip="Similar to the advanced power noise node, positive values increase low frequencies (with colorful effects), negative values increase high frequencies.",
)
.req_float_k(
default=1.0,
tooltip="Currently no description of exactly what it does, it's just another knob you can try turning for a different effect.",
)
.req_float_vertical_factor(
default=1.0,
tooltip="Vertical frequency scaling factor.",
)
.req_float_horizontal_factor(
default=1.0,
tooltip="Horizontal frequency scaling factor.",
)
.req_bool_use_sqrt(
default=True,
tooltip="Controls whether to sqrt when dividing the FFT. Negative hfac/wfac won't work when enabled. Turning it off seems to make the parameters have a much stronger effect.",
)
),
)
@@ -130,31 +134,33 @@ class SonarAdvancedPowerLawNoiseNode(SonarCustomNoiseNodeBase):
DESCRIPTION = "Custom noise type that allows specifying parameters for power law (grey, violet, etc) variants."
INPUT_TYPES = SonarLazyInputTypes(
lambda: NoiseChainInputTypes()
.req_float_alpha(
default=0.5,
tooltip="Similar to the advanced power noise node, positive values increase low frequencies (with colorful effects), negative values increase high frequencies.",
)
.req_field_div_max_dims(
(
"none",
"non-batch",
"spatial",
"all",
"batch",
"channel",
"height",
"width",
),
default="non-batch",
tooltip="If non-none, the noise gets divide by the maximum over this dimension.",
)
.req_bool_use_div_max_abs(
default=True,
tooltip="Only has an effect when div_max_dims is not none. Controls whether maximization is done with the absolute values or raw values.",
)
.req_bool_use_sign(
tooltip="When set, only the sign of the initial noise is used, so -0.5, -0.2 all turn into -1, 0.5, 2, etc all turn into 1.",
lambda: (
NoiseChainInputTypes()
.req_float_alpha(
default=0.5,
tooltip="Similar to the advanced power noise node, positive values increase low frequencies (with colorful effects), negative values increase high frequencies.",
)
.req_field_div_max_dims(
(
"none",
"non-batch",
"spatial",
"all",
"batch",
"channel",
"height",
"width",
),
default="non-batch",
tooltip="If non-none, the noise gets divide by the maximum over this dimension.",
)
.req_bool_use_div_max_abs(
default=True,
tooltip="Only has an effect when div_max_dims is not none. Controls whether maximization is done with the absolute values or raw values.",
)
.req_bool_use_sign(
tooltip="When set, only the sign of the initial noise is used, so -0.5, -0.2 all turn into -1, 0.5, 2, etc all turn into 1.",
)
),
)
@@ -199,114 +205,116 @@ class SonarAdvancedCollatzNoiseNode(SonarCustomNoiseNodeBase):
DESCRIPTION = "Custom noise type that allows specifying parameters for Collatz noise. Very experimental, also very slow. It might just about work as initial noise with non-ancestral sampling but if you get weird results I recommend mixing it with other noise types or possibly using ancestral/SDE sampling."
INPUT_TYPES = SonarLazyInputTypes(
lambda: NoiseChainInputTypes()
.req_bool_adjust_scale(
default=False,
tooltip="When enabled, the output will be normalized to values between -1 and 1 using the last two dimensions (if there are four or more), otherwise dimensions after the first.",
)
.req_string_chain_length(
default="1, 1, 2, 2, 3, 3",
tooltip="Comma-separated list of chain lengths. Cannot be empty. Iterations will cycle through the list and wrap. Controls the length of Collatz chains. Note: Using a high chain length may be very slow, especially if combined with many iterations.",
)
.req_int_chain_offset(
default=5,
min=0,
max=10000,
tooltip="Uses values starting at the specified offset. Note: This entails generating chains of length chain_length + chain_offset, which may be quite slow if you use high values.",
)
.req_int_iterations(
default=10,
min=1,
max=10000,
tooltip="Number of iterations to run. Warning: Collatz noise (my implementation, anyway) is EXTREMELY slow.",
)
.req_bool_iteration_sign_flipping(
default=True,
tooltip="Controls whether we cycle between flipping the sign on the output from each iteration. May average out weirdness... Or make stuff weirder.",
)
.req_float_rmin(
default=-8000.0,
tooltip="Minimum value a chain can start with. Going as low as -9500 should be safe with float32.",
)
.req_float_rmax(
default=8000.0,
tooltip="Maximum value a chain can start with. I don't recommend going over 9500 if you are using the float32 dtype here as that is where the Collatz chain starts to reach values that can't be accurately represented.",
)
.req_string_dims(
default="-1, -1, -2, -2",
tooltip="Comma-separated list of dimensions. Cannot be empty. May be negative to count from the end of the list. Iterations will cycle through the list and wrap.",
)
.req_bool_flatten(
tooltip="Controls whether dimensions past the current one selected from the dims parameter will get flattened.",
)
.req_field_output_mode(
(
"values",
"ratios",
"mults",
"adds",
"seed_x_mults",
"seed_x_adds",
"noise_x_ratios",
"noise_x_mults",
"noise_x_adds",
),
default="values",
)
.req_float_quantile(
default=0.5,
min=0.0,
max=1.0,
tooltip="The initial output of each iteration will be run through quantile normalization. Setting the parameter to 0 or 1 will disable quantile normalization.",
)
.req_field_quantile_strategy(
tuple(utils.quantile_handlers.keys()),
default="clamp",
tooltip="Determines how to treat outliers. zero and reverse_zero modes are only useful if you're going to do something like add the result to some other noise. zero will return zero for anything outside the quantile range, reverse_zero only _keeps_ the outliers and zeros everything else.",
)
.req_field_noise_dtype(
("float32", "float64", "float16", "bfloat16"),
default="float32",
tooltip="Generally should be left at the default. Only float32 and float64 will work if you have quantile normalization enabled.",
)
.req_float_even_multiplier(
default=0.5,
tooltip="Multiplier to use when the previous link in the chain is even. Collatz uses 0.5 (divides by two) here.",
)
.req_float_even_addition(
default=0.0,
tooltip="Value to add when the previous link in the chain is even. Collatz uses 0 here.",
)
.req_float_odd_multiplier(
default=3.0,
tooltip="Multiplier to use when the previous link in the chain is odd. Collatz uses 3 here.",
)
.req_float_odd_addition(
default=1.0,
tooltip="Value to add when the previous link in the chain is odd. Collatz uses 1 here.",
)
.req_bool_integer_math(
default=True,
tooltip="Controls whether the results during chain generation get truncated to an integer value or not. Should be enabled if you actually want to generate accurate Collatz chains.",
)
.req_bool_add_preserves_sign(
default=True,
tooltip="Controls whether additions use the same sign as the item they're being added to.",
)
.req_bool_break_loops(
default=True,
tooltip="Controls whether the chain resets back to the seed value once it reaches 1 or 0. Generally should be left enabled, otherwise the chain will oscillate between only a few values for the rest of the length (at least with the Collatz rules).",
)
.req_field_seed_mode(
("default", "force_odd", "force_even"),
default="default",
tooltip="Default mode just uses whatever the original seed value was. force_odd/force_even will force it to the specified parity by adding one if it doesn't match. Starting from odd seeds might result in longer chains. Enabling the force modes may cause the initial seeds to exceed rmax by one.",
)
.opt_customnoise_seed_custom_noise(
tooltip="Optional custom noise to use for initial values for Collatz chains. May be slow as it will generate noise according to the original input size and then crop it. Does this noise type have enough warnings about it being slow? Yeah. Connecting something here will probably make it even slower!",
)
.opt_customnoise_mix_custom_noise(
tooltip="Optional custom noise to use with the output modes starting with 'noise'.",
lambda: (
NoiseChainInputTypes()
.req_bool_adjust_scale(
default=False,
tooltip="When enabled, the output will be normalized to values between -1 and 1 using the last two dimensions (if there are four or more), otherwise dimensions after the first.",
)
.req_string_chain_length(
default="1, 1, 2, 2, 3, 3",
tooltip="Comma-separated list of chain lengths. Cannot be empty. Iterations will cycle through the list and wrap. Controls the length of Collatz chains. Note: Using a high chain length may be very slow, especially if combined with many iterations.",
)
.req_int_chain_offset(
default=5,
min=0,
max=10000,
tooltip="Uses values starting at the specified offset. Note: This entails generating chains of length chain_length + chain_offset, which may be quite slow if you use high values.",
)
.req_int_iterations(
default=10,
min=1,
max=10000,
tooltip="Number of iterations to run. Warning: Collatz noise (my implementation, anyway) is EXTREMELY slow.",
)
.req_bool_iteration_sign_flipping(
default=True,
tooltip="Controls whether we cycle between flipping the sign on the output from each iteration. May average out weirdness... Or make stuff weirder.",
)
.req_float_rmin(
default=-8000.0,
tooltip="Minimum value a chain can start with. Going as low as -9500 should be safe with float32.",
)
.req_float_rmax(
default=8000.0,
tooltip="Maximum value a chain can start with. I don't recommend going over 9500 if you are using the float32 dtype here as that is where the Collatz chain starts to reach values that can't be accurately represented.",
)
.req_string_dims(
default="-1, -1, -2, -2",
tooltip="Comma-separated list of dimensions. Cannot be empty. May be negative to count from the end of the list. Iterations will cycle through the list and wrap.",
)
.req_bool_flatten(
tooltip="Controls whether dimensions past the current one selected from the dims parameter will get flattened.",
)
.req_field_output_mode(
(
"values",
"ratios",
"mults",
"adds",
"seed_x_mults",
"seed_x_adds",
"noise_x_ratios",
"noise_x_mults",
"noise_x_adds",
),
default="values",
)
.req_float_quantile(
default=0.5,
min=0.0,
max=1.0,
tooltip="The initial output of each iteration will be run through quantile normalization. Setting the parameter to 0 or 1 will disable quantile normalization.",
)
.req_field_quantile_strategy(
tuple(utils.quantile_handlers.keys()),
default="clamp",
tooltip="Determines how to treat outliers. zero and reverse_zero modes are only useful if you're going to do something like add the result to some other noise. zero will return zero for anything outside the quantile range, reverse_zero only _keeps_ the outliers and zeros everything else.",
)
.req_field_noise_dtype(
("float32", "float64", "float16", "bfloat16"),
default="float32",
tooltip="Generally should be left at the default. Only float32 and float64 will work if you have quantile normalization enabled.",
)
.req_float_even_multiplier(
default=0.5,
tooltip="Multiplier to use when the previous link in the chain is even. Collatz uses 0.5 (divides by two) here.",
)
.req_float_even_addition(
default=0.0,
tooltip="Value to add when the previous link in the chain is even. Collatz uses 0 here.",
)
.req_float_odd_multiplier(
default=3.0,
tooltip="Multiplier to use when the previous link in the chain is odd. Collatz uses 3 here.",
)
.req_float_odd_addition(
default=1.0,
tooltip="Value to add when the previous link in the chain is odd. Collatz uses 1 here.",
)
.req_bool_integer_math(
default=True,
tooltip="Controls whether the results during chain generation get truncated to an integer value or not. Should be enabled if you actually want to generate accurate Collatz chains.",
)
.req_bool_add_preserves_sign(
default=True,
tooltip="Controls whether additions use the same sign as the item they're being added to.",
)
.req_bool_break_loops(
default=True,
tooltip="Controls whether the chain resets back to the seed value once it reaches 1 or 0. Generally should be left enabled, otherwise the chain will oscillate between only a few values for the rest of the length (at least with the Collatz rules).",
)
.req_field_seed_mode(
("default", "force_odd", "force_even"),
default="default",
tooltip="Default mode just uses whatever the original seed value was. force_odd/force_even will force it to the specified parity by adding one if it doesn't match. Starting from odd seeds might result in longer chains. Enabling the force modes may cause the initial seeds to exceed rmax by one.",
)
.opt_customnoise_seed_custom_noise(
tooltip="Optional custom noise to use for initial values for Collatz chains. May be slow as it will generate noise according to the original input size and then crop it. Does this noise type have enough warnings about it being slow? Yeah. Connecting something here will probably make it even slower!",
)
.opt_customnoise_mix_custom_noise(
tooltip="Optional custom noise to use with the output modes starting with 'noise'.",
)
),
)
@@ -486,68 +494,70 @@ class SonarWaveletNoiseNode(
DESCRIPTION = "Custom noise type that allows generating wavelet noise. Very simple explanation of how a single octave works:\n1) Generate some noise.\n2) Scale it down 50%.\n3) Scale it back up to the original size.\n4) Subtract the scaled noise from the original noise.\nScaling the noise down and then back up blurs it, so this is essentially sharpening the noise."
INPUT_TYPES = SonarLazyInputTypes(
lambda: NoiseChainInputTypes()
.req_int_octaves(
default=4,
min=-100,
max=100,
tooltip="Number of octaves to generate. You can use a negative number here to run the octaves in reverse order though it may produce weird results/not work very well.",
)
.req_float_octave_height_factor(
default=0.5,
min=0.001,
tooltip="Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.",
)
.req_float_octave_width_factor(
default=0.5,
min=0.001,
tooltip="Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.",
)
.req_selectscalemode_octave_scale_mode(
default="adaptive_avg_pool2d",
tooltip="Scaling mode used within each octave to produce the scaled noise. By default this will be scaling down that octave's noise.",
)
.req_selectscalemode_octave_rescale_mode(
default="bilinear",
tooltip="Scaling mode used within each octave to scale the noise back up to that octave's original size.",
)
.req_selectscalemode_post_octave_rescale_mode(
default="bilinear",
tooltip="Scaling mode used to scale the output of an octave back up to the actual latent size.",
)
.req_float_initial_amplitude(
default=1.0,
tooltip="Basically the strength an octave gets added to the total. This will be scaled by persistance after each octave.",
)
.req_float_persistence(
default=0.5,
tooltip="Multiplier applied to amplitude after each octave. 0.5 means the first octave uses initial_amplitude, the second uses half of that and so on.",
)
.req_float_height_factor(
default=2.0,
min=0.001,
tooltip="Scaling factor for height, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.",
)
.req_float_width_factor(
tooltip="Scaling factor for width, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.",
default=2.0,
min=0.001,
)
.req_float_update_blend(
tooltip="Controls how original_noise - scaled_noise is blended with original_noise. The default is to use 100% original_noise - scaled_noise.",
default=1.0,
)
.req_selectblend_update_blend_mode(
insert_modes=("simple_add",),
default="lerp",
tooltip="Controls how the enhanced noise from each octave is blended with that octave's raw noise. With normal wavelet noise there's no blending and you use 100% enhanced noise.",
)
.req_bool_normalize_noise(
tooltip="Controls whether the noise source is normalized before wavelet filtering occurs.",
)
.req_normalizetristate_normalize()
.opt_customnoise_custom_noise(
tooltip="Optional: Custom noise input. If unconnected will default to Gaussian noise. Note: When connected, the noise for all octaves will be generated at the maximum scale and then cropped which may be slow.",
lambda: (
NoiseChainInputTypes()
.req_int_octaves(
default=4,
min=-100,
max=100,
tooltip="Number of octaves to generate. You can use a negative number here to run the octaves in reverse order though it may produce weird results/not work very well.",
)
.req_float_octave_height_factor(
default=0.5,
min=0.001,
tooltip="Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.",
)
.req_float_octave_width_factor(
default=0.5,
min=0.001,
tooltip="Wavelet noise works by scaling noise by this factor in each octave, then scaling it back up to the original size. After that, the scaled noise is subtracted from the original noise.",
)
.req_selectscalemode_octave_scale_mode(
default="adaptive_avg_pool2d",
tooltip="Scaling mode used within each octave to produce the scaled noise. By default this will be scaling down that octave's noise.",
)
.req_selectscalemode_octave_rescale_mode(
default="bilinear",
tooltip="Scaling mode used within each octave to scale the noise back up to that octave's original size.",
)
.req_selectscalemode_post_octave_rescale_mode(
default="bilinear",
tooltip="Scaling mode used to scale the output of an octave back up to the actual latent size.",
)
.req_float_initial_amplitude(
default=1.0,
tooltip="Basically the strength an octave gets added to the total. This will be scaled by persistance after each octave.",
)
.req_float_persistence(
default=0.5,
tooltip="Multiplier applied to amplitude after each octave. 0.5 means the first octave uses initial_amplitude, the second uses half of that and so on.",
)
.req_float_height_factor(
default=2.0,
min=0.001,
tooltip="Scaling factor for height, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.",
)
.req_float_width_factor(
tooltip="Scaling factor for width, calculated after each octave. 2.0 means divide by two. Note: It's possible to use values below 1 here but be careful as it's very easy to reach absurd latent sizes with only a few octaves.",
default=2.0,
min=0.001,
)
.req_float_update_blend(
tooltip="Controls how original_noise - scaled_noise is blended with original_noise. The default is to use 100% original_noise - scaled_noise.",
default=1.0,
)
.req_selectblend_update_blend_mode(
insert_modes=("simple_add",),
default="lerp",
tooltip="Controls how the enhanced noise from each octave is blended with that octave's raw noise. With normal wavelet noise there's no blending and you use 100% enhanced noise.",
)
.req_bool_normalize_noise(
tooltip="Controls whether the noise source is normalized before wavelet filtering occurs.",
)
.req_normalizetristate_normalize()
.opt_customnoise_custom_noise(
tooltip="Optional: Custom noise input. If unconnected will default to Gaussian noise. Note: When connected, the noise for all octaves will be generated at the maximum scale and then cropped which may be slow.",
)
),
)
@@ -612,77 +622,79 @@ class SonarAdvancedVoronoiNoiseNode(SonarCustomNoiseNodeBase):
),
_pretty_result_modes=", ".join( # noqa: B008
sorted(VoronoiNoiseGenerator.voronoi_result_modes), # noqa: B008
): NoiseChainInputTypes()
.req_string_n_points(
default="256",
tooltip="Controls the number of features points in the generated noise. Higher generally results in more detail/better results but is slower. May be a comma separated list for each octave (only applicable when octave mode is set to new_features). 2 is the minimum value.",
)
.req_string_distance_mode(
default="euclidean",
placeholder=f"One of: {_pretty_distance_modes}",
tooltip="Distance modes. You can specify a comma-separated list of items which will be used for each octave.\n"
"You can specify an average of multiple distance modes by separating the names with +.\n"
"Some modes can take arguments. Example syntax: modename:argname=value:argname=value\n"
"All modes support scaling their output with dscale (which defaults to 1).\n"
f"Possible distance modes: {_pretty_distance_modes}",
)
.req_float_z_initial(
default=0.0,
tooltip="Initial value for z (depth).",
)
.req_float_z_increment(
default=1.0,
tooltip="Amount z (depth) is incremented when applicable.",
)
.req_float_z_max(
default=9999.0,
tooltip="Maximum difference from the intial value. At that point, z_max_mode will apply. When set to 0, z_increment has no effect and you will get different noise each time you call the noise sampler.",
)
.req_field_z_max_mode(
(
"reset",
"wrap",
"bounce",
),
default="reset",
tooltip="Controls what happens when the z_max limit is hit (see tooltip for z_max). Reset will reset the feature points and z to the initial values. Wrap will reset z to the initial value. Bounce will flip the sign on the increment and do an increment.",
)
.req_string_result_mode(
default="diff2",
placeholder=f"One of: {_pretty_result_modes}",
tooltip="Result modes. You can specify a comma-separated list of items which will be used for each octave.\n"
"You can specify an average of multiple result modes by separating the names with +.\n"
"Some modes can take arguments. Example syntax: modename:argname=value:argname=value\n"
"All modes support scaling their output with rscale (which defaults to 1).\n"
f"Possible result modes: {_pretty_result_modes}",
)
.req_field_octave_mode(
(
"same_features",
"new_features",
"same_invert_odd",
"same_invert_even",
"same_roll_chan_up",
"same_roll_chan_down",
"same_roll_dir_up",
"same_roll_dir_down",
),
default="new_features",
tooltip="Only relevant when generating multiple octaves. Controls whether octaves share a set of feature points or if they are different for each octave (note that this is slower). Modes starting with 'same' will use the same feature points per octave but may transform them.",
)
.req_int_octaves(
default=3,
min=1,
tooltip="Number of octaves of noise to generate.",
)
.req_float_gain(default=0.75)
.req_float_lacunarity(default=2.0)
.req_float_initial_amplitude(default=1.0)
.req_float_initial_scale(default=1.0)
.req_normalizetristate_normalize()
.opt_customnoise(
"custom_noise",
tooltip="Optional input if you want to use some other noise type for the initial feature points. Won't work well with noise types that care about the content of the latent (I think only spectral modulation) or manage their own seed (I believe this only applies to Brownian or if you're using the custom noise parameters node to override seeds/fork the RNG).",
): (
NoiseChainInputTypes()
.req_string_n_points(
default="256",
tooltip="Controls the number of features points in the generated noise. Higher generally results in more detail/better results but is slower. May be a comma separated list for each octave (only applicable when octave mode is set to new_features). 2 is the minimum value.",
)
.req_string_distance_mode(
default="euclidean",
placeholder=f"One of: {_pretty_distance_modes}",
tooltip="Distance modes. You can specify a comma-separated list of items which will be used for each octave.\n"
"You can specify an average of multiple distance modes by separating the names with +.\n"
"Some modes can take arguments. Example syntax: modename:argname=value:argname=value\n"
"All modes support scaling their output with dscale (which defaults to 1).\n"
f"Possible distance modes: {_pretty_distance_modes}",
)
.req_float_z_initial(
default=0.0,
tooltip="Initial value for z (depth).",
)
.req_float_z_increment(
default=1.0,
tooltip="Amount z (depth) is incremented when applicable.",
)
.req_float_z_max(
default=9999.0,
tooltip="Maximum difference from the intial value. At that point, z_max_mode will apply. When set to 0, z_increment has no effect and you will get different noise each time you call the noise sampler.",
)
.req_field_z_max_mode(
(
"reset",
"wrap",
"bounce",
),
default="reset",
tooltip="Controls what happens when the z_max limit is hit (see tooltip for z_max). Reset will reset the feature points and z to the initial values. Wrap will reset z to the initial value. Bounce will flip the sign on the increment and do an increment.",
)
.req_string_result_mode(
default="diff2",
placeholder=f"One of: {_pretty_result_modes}",
tooltip="Result modes. You can specify a comma-separated list of items which will be used for each octave.\n"
"You can specify an average of multiple result modes by separating the names with +.\n"
"Some modes can take arguments. Example syntax: modename:argname=value:argname=value\n"
"All modes support scaling their output with rscale (which defaults to 1).\n"
f"Possible result modes: {_pretty_result_modes}",
)
.req_field_octave_mode(
(
"same_features",
"new_features",
"same_invert_odd",
"same_invert_even",
"same_roll_chan_up",
"same_roll_chan_down",
"same_roll_dir_up",
"same_roll_dir_down",
),
default="new_features",
tooltip="Only relevant when generating multiple octaves. Controls whether octaves share a set of feature points or if they are different for each octave (note that this is slower). Modes starting with 'same' will use the same feature points per octave but may transform them.",
)
.req_int_octaves(
default=3,
min=1,
tooltip="Number of octaves of noise to generate.",
)
.req_float_gain(default=0.75)
.req_float_lacunarity(default=2.0)
.req_float_initial_amplitude(default=1.0)
.req_float_initial_scale(default=1.0)
.req_normalizetristate_normalize()
.opt_customnoise(
"custom_noise",
tooltip="Optional input if you want to use some other noise type for the initial feature points. Won't work well with noise types that care about the content of the latent (I think only spectral modulation) or manage their own seed (I believe this only applies to Brownian or if you're using the custom noise parameters node to override seeds/fork the RNG).",
)
),
)
@@ -737,12 +749,271 @@ class SonarAdvancedVoronoiNoiseNode(SonarCustomNoiseNodeBase):
)
class SonarAdvancedSimulationNoiseNode(SonarCustomNoiseNodeBase):
DESCRIPTION = "TBD"
INPUT_TYPES = SonarLazyInputTypes(
lambda: (
NoiseChainInputTypes()
.req_field_spectral_mode(
("multi_octave", "power_law", "band_pass"),
default="multi_octave",
tooltip="TBD",
)
.req_field_field_mode(
("basis", "curl", "projection", "curl_ndim", "basis_ndim"),
default="basis",
tooltip="TBD",
)
.req_field_band_shape(
("log_gaussian", "raised_cosine"),
default="log_gaussian",
tooltip="No effect in power_law spectral mode or in multi_octave spectral mode when octaves is set to 0.",
)
.req_field_channel_mode(
(
"stacked",
"over_depth",
"flat",
"over_depth_alt",
"over_depth_avg",
"over_depth_h",
"over_depth_w",
"over_depth_z",
"over_depth_h_sub_w",
"over_depth_h_sub_z",
"over_depth_w_sub_h",
"over_depth_w_sub_z",
"over_depth_z_sub_h",
"over_depth_z_sub_w",
"over_depth_h_add_w",
"over_depth_h_add_z",
"over_depth_w_add_h",
"over_depth_w_add_z",
"over_depth_z_add_h",
"over_depth_z_add_w",
"over_depth_h_mul_w",
"over_depth_h_mul_z",
"over_depth_w_mul_h",
"over_depth_w_mul_z",
"over_depth_z_mul_h",
"over_depth_z_mul_w",
"over_depth_h_div_w",
"over_depth_h_div_z",
"over_depth_w_div_h",
"over_depth_w_div_z",
"over_depth_z_div_h",
"over_depth_z_div_w",
),
default="over_depth",
tooltip="TBD",
)
.req_int_depth(default=16)
.req_int_initial_depth(
default=0,
min=0,
tooltip="TBD",
)
.req_int_max_depth(
default=-1,
tooltip="TBD",
)
.req_field_depth_mode(
("reset", "wrap", "bounce"),
default="reset",
tooltip="TBD",
)
.req_int_octaves(
default=3,
min=0,
tooltip="Number of octaves of noise to generate. Only has an effect in multi_octave spectral mode. You can also set octaves to 0 to disable octaves.",
)
.req_float_lacunarity(default=2.0)
.req_float_gain(default=0.75)
.req_float_log_gaussian_sigma(default=0.3)
.req_string_anisotropy(
default="",
tooltip="TBD",
)
.req_normalizetristate_normalize()
.req_float_base_k(default=0.0)
.req_float_power_law_beta(
default=0.25,
tooltip="Beta used for power_law spectral mode, no effect otherwise. Higher beta will result in colorful low frequency noise, low (or negative) will emphasize high frequencies.",
)
.req_float_band_pass_low(default=0.0001, min=1e-06)
.req_float_band_pass_high(default=1.0, min=1e-06)
.opt_customnoise_custom_noise_h(
tooltip="TBD",
)
.opt_customnoise_custom_noise_w(
tooltip="TBD",
)
.opt_customnoise_custom_noise_z(
tooltip="TBD",
)
),
)
@classmethod
def get_item_class(cls):
return noise.AdvancedSimulationNoise
def go(
self,
*,
factor: float,
rescale: float,
depth: int,
initial_depth: int,
max_depth: int,
depth_mode: str,
octaves: int,
gain: float,
lacunarity: float,
channel_mode: str,
band_shape: str,
log_gaussian_sigma: float,
anisotropy: str,
normalize: str,
spectral_mode: str,
field_mode: str,
base_k: float,
power_law_beta: float,
band_pass_low: float,
band_pass_high: float,
sonar_custom_noise_opt=None,
custom_noise_h=None,
custom_noise_w=None,
custom_noise_z=None,
):
anisotropy = anisotropy.strip()
anisotropy = (
None
if not anisotropy
else tuple(float(v) if v.strip() else 1.0 for v in anisotropy.split(","))
)
return super().go(
factor,
rescale=rescale,
sonar_custom_noise_opt=sonar_custom_noise_opt,
depth=depth,
initial_depth=initial_depth,
max_depth=max_depth,
depth_mode=depth_mode,
octaves=octaves,
gain=gain,
lacunarity=lacunarity,
channel_mode=channel_mode,
band_shape=band_shape,
log_gaussian_sigma=log_gaussian_sigma,
anisotropy=anisotropy,
spectral_mode=spectral_mode,
field_mode=field_mode,
base_k=base_k,
power_law_beta=power_law_beta,
band_pass_low=band_pass_low,
band_pass_high=band_pass_high,
normalize=normalize,
custom_noise_h=custom_noise_h.clone()
if custom_noise_h is not None
else None,
custom_noise_w=custom_noise_w.clone()
if custom_noise_w is not None
else None,
custom_noise_z=custom_noise_z.clone()
if custom_noise_z is not None
else None,
)
class SonarAdvancedAutomataNoiseNode(SonarCustomNoiseNodeBase):
DESCRIPTION = "TBD"
INPUT_TYPES = SonarLazyInputTypes(
lambda: (
NoiseChainInputTypes()
.req_int_num_seeds(
default=20,
min=1,
tooltip="Controls the initial number of seeds.",
)
.req_int_depth(
default=10,
min=0,
tooltip="Controls the depth. Set to 0 to disable 3D noise generation. Note: The whole 3D chunk has to be generated at once which may be memory intensive.",
)
.req_int_steps(
default=20,
min=0,
tooltip="Controls the initial number of seeds.",
)
.req_int_spread_substeps(
default=2,
min=0,
tooltip="TBD",
)
.req_field_evolution_mode(
("collatz",),
default="collatz",
tooltip="Contols the main noise evolution function.",
)
.req_field_spread_mode(
("blur",),
default="blur",
tooltip="Controls how the values spread each noise step. Blur uses convolution.",
)
.req_normalizetristate_normalize()
.opt_customnoise(
"custom_noise",
tooltip="Optional input if you want to use some other noise type for the initial feature points. Won't work well with noise types that care about the content of the latent (I think only spectral modulation) or manage their own seed (I believe this only applies to Brownian or if you're using the custom noise parameters node to override seeds/fork the RNG).",
)
),
)
@classmethod
def get_item_class(cls):
return noise.AdvancedAutomataNoise
def go(
self,
*,
factor: float,
rescale: float,
num_seeds: int,
depth: int,
steps: int,
spread_substeps: int,
evolution_mode: str,
spread_mode: str,
normalize: str,
custom_noise=None,
sonar_custom_noise_opt=None,
): # ty:ignore[invalid-method-override]
return super().go(
factor,
rescale=rescale,
sonar_custom_noise_opt=sonar_custom_noise_opt,
num_seeds=num_seeds,
depth=depth,
steps=steps,
spread_substeps=spread_substeps,
evolution_mode=evolution_mode,
spread_mode=spread_mode,
custom_noise=custom_noise,
normalize=normalize,
)
NODE_CLASS_MAPPINGS = {
"SonarAdvancedPyramidNoise": SonarAdvancedPyramidNoiseNode,
"SonarAdvanced1fNoise": SonarAdvanced1fNoiseNode,
"SonarAdvancedPowerLawNoise": SonarAdvancedPowerLawNoiseNode,
"SonarAdvancedAutomataNoise": SonarAdvancedAutomataNoiseNode,
"SonarAdvancedCollatzNoise": SonarAdvancedCollatzNoiseNode,
"SonarAdvancedDistroNoise": SonarAdvancedDistroNoiseNode,
"SonarAdvancedPowerLawNoise": SonarAdvancedPowerLawNoiseNode,
"SonarAdvancedPyramidNoise": SonarAdvancedPyramidNoiseNode,
"SonarAdvancedSimulationNoise": SonarAdvancedSimulationNoiseNode,
"SonarAdvancedVoronoiNoise": SonarAdvancedVoronoiNoiseNode,
"SonarWaveletNoise": SonarWaveletNoiseNode,
}
+7 -6
View File
@@ -7,6 +7,7 @@ from __future__ import annotations
import math
import os
import random
from typing import Any
import comfy
import folder_paths
@@ -86,7 +87,7 @@ class ChannelMixer:
channel_mixer /= channel_mixer.norm(dim=1, keepdim=True)
return channel_mixer
def to(self, *args: list, **kwargs: dict):
def to(self, *args: Any, **kwargs: Any):
if self.mixer is not None:
self.mixer = self.mixer.to(*args, **kwargs)
return self
@@ -100,7 +101,7 @@ class ChannelMixer:
noise = self.mixer @ noise.swapaxes(0, 1).reshape(c, -1)
return noise.reshape(c, b, h, w).swapaxes(1, 0)
def __call__(self, *args: list, **kwargs: dict):
def __call__(self, *args: Any, **kwargs: Any):
return self.apply(*args, **kwargs)
@@ -301,7 +302,7 @@ class PowerNoiseItem(CustomNoiseItemBase):
*,
channel_correlation,
power_filter=None,
**kwargs: dict,
**kwargs: Any,
):
if isinstance(channel_correlation, str):
channel_correlation = torch.tensor(
@@ -476,7 +477,7 @@ class PowerFilterNoiseItem(PowerNoiseItem):
noise,
normalize_noise,
normalize_result,
**kwargs: dict,
**kwargs: Any,
):
super().__init__(
factor,
@@ -628,7 +629,7 @@ class SonarPowerNoiseNode(SonarCustomNoiseNodeBase):
def go(
self,
preview="none",
**kwargs: dict,
**kwargs: Any,
):
result = super().go(**kwargs)
if preview == "none":
@@ -713,7 +714,7 @@ class SonarPowerFilterNoiseNode(SonarPowerNoiseNode, SonarNormalizeNoiseNodeMixi
normalize_noise,
normalize_result,
preview="none",
**kwargs: dict,
**kwargs: Any,
):
return super().go(
factor=factor,
+273 -73
View File
@@ -4,12 +4,14 @@ import abc
import math
import random
from functools import partial
from typing import Callable
from typing import TYPE_CHECKING
import comfy
import torch
import yaml
from comfy import utils as comfy_utils
from comfy.k_diffusion import sampling
from comfy.model_management import throw_exception_if_processing_interrupted
from torch import Tensor
from . import external, utils
@@ -24,6 +26,14 @@ from .utils import (
scale_noise,
)
try:
from comfy import nested_tensor
except (ModuleNotFoundError, ImportError):
nested_tensor = None
if TYPE_CHECKING:
from collections.abc import Callable
# ruff: noqa: ANN002, ANN003, FBT001
@@ -64,7 +74,16 @@ class CustomNoiseItemBase(abc.ABC):
def get_normalize(self, k, default=None):
val = getattr(self, k, None)
return default if val is None else val
if val in {None, "default"}:
return default
if val == "disabled":
return False
if val == "forced":
return True
return default
# if isinstance(val, bool):
# return val
# return default if val is None else val
@abc.abstractmethod
def make_noise_sampler(
@@ -456,7 +475,7 @@ class AdvancedVoronoiNoise(AdvancedNoiseBase):
return super().clone_key(k)
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
if x.ndim != 4:
if x.ndim < 4:
raise ValueError("Can only handle 4+ dimensional latents")
return super().make_noise_sampler(
x,
@@ -467,6 +486,59 @@ class AdvancedVoronoiNoise(AdvancedNoiseBase):
)
class AdvancedSimulationNoise(AdvancedNoiseBase):
ns_factory_arg_keys = tuple(SimulationNoiseGenerator.ng_params(no_super=True))
@property
def ns_factory(self):
return SimulationNoiseGenerator
def clone_key(self, k):
if (
k in {"custom_noise_h", "custom_noise_w", "custom_noise_z"}
and getattr(self, k) is not None
):
return getattr(self, k).clone()
return super().clone_key(k)
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
if x.ndim < 4:
raise ValueError("Can only handle 4+ dimensional latents")
return super().make_noise_sampler(
x,
*args,
normalized=normalized,
noise_sampler_factory_h=self.custom_noise_h,
noise_sampler_factory_w=self.custom_noise_w,
noise_sampler_factory_z=self.custom_noise_z,
**kwargs,
)
class AdvancedAutomataNoise(AdvancedNoiseBase):
ns_factory_arg_keys = tuple(AutomataNoiseGenerator.ng_params())
@property
def ns_factory(self):
return AutomataNoiseGenerator
# def clone_key(self, k):
# if k == "custom_noise" and self.custom_noise is not None:
# return self.custom_noise.clone()
# return super().clone_key(k)
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
if x.ndim < 4:
raise ValueError("Can only handle 4+ dimensional latents")
return super().make_noise_sampler(
x,
*args,
normalized=normalized,
# noise_sampler_factory=self.custom_noise,
**kwargs,
)
class CompositeNoise(CustomNoiseItemBase):
def __init__(
self,
@@ -1311,17 +1383,9 @@ class BlendedNoise(CustomNoiseItemBase):
custom_noise_mask=None,
noise_2_percent=0.5,
):
if custom_noise_1 is None and (
custom_noise_mask is not None or noise_2_percent != 1
):
if custom_noise_1 is None and custom_noise_2 is None:
raise ValueError(
"When custom_noise_1 is not attached noise_2_percent must be set to 1",
)
if custom_noise_2 is None and (
custom_noise_mask is not None or noise_2_percent != 0
):
raise ValueError(
"When custom_noise_2 is not attached noise_2_percent must be set to 0",
"At least one of the custom_noise inputs must be connected.",
)
if (
custom_noise_mask is None
@@ -1334,7 +1398,7 @@ class BlendedNoise(CustomNoiseItemBase):
factor,
noise_2_percent=noise_2_percent,
blend_function=blend_function,
custom_noise_1=custom_noise_1.clone(),
custom_noise_1=None if custom_noise_1 is None else custom_noise_1.clone(),
custom_noise_2=None if custom_noise_2 is None else custom_noise_2.clone(),
custom_noise_mask=None
if custom_noise_mask is None
@@ -1344,7 +1408,7 @@ class BlendedNoise(CustomNoiseItemBase):
def clone_key(self, k):
if k == "custom_noise_1":
return self.custom_noise_1.clone()
return None if self.custom_noise_1 is None else self.custom_noise_1.clone()
if k == "custom_noise_2":
return None if self.custom_noise_2 is None else self.custom_noise_2.clone()
if k == "custom_noise_mask":
@@ -1361,47 +1425,33 @@ class BlendedNoise(CustomNoiseItemBase):
blend_function = self.blend_function
n2_blend = self.noise_2_percent
ns_1 = self.custom_noise_1.make_noise_sampler(
x,
*args,
normalized=False,
**kwargs,
)
ns_2 = (
if self.custom_noise_1 is None and self.custom_noise_2 is None:
raise RuntimeError("Impossible: No available noise generator")
ns_1, ns_2, ns_mask = (
None
if self.custom_noise_2 is None
else self.custom_noise_2.make_noise_sampler(
x,
*args,
normalized=False,
**kwargs,
)
)
ns_mask = (
None
if self.custom_noise_mask is None
else self.custom_noise_mask.make_noise_sampler(
x,
*args,
normalized=False,
**kwargs,
)
if ng is None
else ng.make_noise_sampler(x, *args, normalized=False, **kwargs)
for ng in (self.custom_noise_1, self.custom_noise_2, self.custom_noise_mask)
)
n2_blend_tensor = x.new_full((1,), n2_blend) if ns_mask is None else None
def noise_sampler(s, sn):
def noise_sampler(s, sn, *args, **kwargs):
nonlocal n2_blend_tensor
noise_1 = ns_1(s, sn)
noise_2 = None if ns_2 is None else ns_2(s, sn)
noise_1 = None if ns_1 is None else ns_1(s, sn, *args, **kwargs)
noise_2 = None if ns_2 is None else ns_2(s, sn, *args, **kwargs)
if noise_1 is None:
noise_1 = noise_2
elif noise_2 is None:
noise_2 = noise_1
if ns_mask is not None:
n2_blend_tensor = (
utils.normalize_to_scale(ns_mask(s, sn), 0.0, 1.0) + n2_blend
utils.normalize_to_scale(
ns_mask(s, sn, *args, **kwargs),
0.0,
1.0,
).add_(n2_blend)
).clamp_(0.0, 1.0)
noise = (
noise_1
if noise_2 is None
else blend_function(noise_1, noise_2, n2_blend_tensor)
)
noise = blend_function(noise_1, noise_2, n2_blend_tensor)
return scale_noise(noise, factor, normalized=normalize)
return noise_sampler
@@ -1662,6 +1712,62 @@ class ScatternetFilteredNoise(CustomNoiseItemBase):
return noise_sampler
class NoveltyFilteredNoise(CustomNoiseItemBase):
def clone_key(self, k):
if k == "noise" and self.noise is not None:
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)
if self.noise is not None:
internal_ns = self.noise.make_noise_sampler(
x,
*args,
sigma_min=sigma_min,
sigma_max=sigma_max,
normalized=self.normalize_noise,
**kwargs,
)
else:
internal_ns = None
ns_kwargs = getattr(self, "ns_kwargs", {}).copy()
kwargs |= ns_kwargs
ns = NoveltyFilteredNoiseGenerator(
x,
*args,
sigma_min=sigma_min,
sigma_max=sigma_max,
normalized=False,
noise_sampler=internal_ns,
skip_initial=self.skip_initial,
iters_per_call=self.iters_per_call,
blend_ratio=self.blend_ratio,
blend_function=utils.BLENDING_MODES[self.blend_mode],
update_blend_ratio=self.update_blend_ratio,
update_blend_function=utils.BLENDING_MODES[self.update_blend_mode],
**kwargs,
)
def noise_sampler(sigma, sigma_next):
return scale_noise(
ns(sigma, sigma_next),
factor,
normalized=normalize,
)
return noise_sampler
class LatentOperationFilteredNoise(CustomNoiseItemBase):
def clone_key(self, k):
if k == "noise" and self.noise is not None:
@@ -1778,6 +1884,10 @@ class QuantileFilteredNoise(CustomNoiseItemBase):
def clone_key(self, k):
if k == "noise":
return self.noise.clone()
if k == "noise_reference":
return (
None if self.noise_reference is None else self.noise_reference.clone()
)
return super().clone_key(k)
def make_noise_sampler(
@@ -1799,6 +1909,18 @@ class QuantileFilteredNoise(CustomNoiseItemBase):
normalized=self.normalize_noise,
**kwargs,
)
ns_ref = (
None
if self.noise_reference is None or self.nq_lo is not None
else self.noise_reference.make_noise_sampler(
x,
*args,
sigma_min=sigma_min,
sigma_max=sigma_max,
normalized=self.normalize_noise,
**kwargs,
)
)
noise_filter = partial(
quantile_normalize,
quantile=self.quantile,
@@ -1807,11 +1929,22 @@ class QuantileFilteredNoise(CustomNoiseItemBase):
nq_fac=self.norm_fac,
pow_fac=self.norm_pow,
strategy=self.strategy,
pow_fac_in=self.norm_power_in,
sign_mode=self.sign_mode,
abs_quantiles=self.abs_quantiles,
only_outliers=self.only_outliers,
nq_lo=self.nq_lo,
nq_hi=self.nq_hi,
zero_mean_scale=self.zero_mean_scale,
restore_mean_scale=self.restore_mean_scale,
zero_result_mean_scale=self.zero_result_mean_scale,
)
def noise_sampler(sigma, sigma_next):
def noise_sampler(*args, **kwargs):
noise_reference = None if ns_ref is None else ns_ref(*args, **kwargs)
noise = ns(*args, **kwargs)
return scale_noise(
noise_filter(ns(sigma, sigma_next)),
noise_filter(noise, noise_reference=noise_reference),
factor,
normalized=normalize,
)
@@ -1866,29 +1999,32 @@ class PerDimNoise(CustomNoiseItemBase):
slice(-dim_size, None) if d == dim else slice(None, None)
for d in range(x.ndim)
)
n_chunks = math.ceil(dim_size / chunk_size)
if self.shrink_dim:
def noise_sampler(sigma, sigma_next) -> torch.Tensor:
noise = torch.cat(
tuple(ns(sigma, sigma_next) for _ in range(dim_size)),
dim=dim,
)[trim_slice]
chunks = []
for _ in range(n_chunks):
throw_exception_if_processing_interrupted()
chunks.append(ns(sigma, sigma_next))
noise = torch.cat(chunks, dim=dim)[trim_slice]
return scale_noise(noise, factor, normalized=normalize)
else:
select_dim = [slice(None, None) for d in range(x.ndim)]
n_chunks = math.ceil(dim_size / chunk_size)
temp_shape = list(x.shape)
temp_shape[dim] = int(n_chunks * chunk_size)
return noise_sampler
def noise_sampler(sigma, sigma_next) -> torch.Tensor:
nonlocal select_dim
result = x.new_zeros(temp_shape)
# result = torch.zeros_like(x)
for idx in range(0, dim_size, chunk_size):
select_dim[dim] = slice(idx, idx + chunk_size)
result[select_dim] = ns(sigma, sigma_next)[select_dim]
return scale_noise(result[trim_slice], factor, normalized=normalize)
select_dim = [slice(None, None) for d in range(x.ndim)]
temp_shape = list(x.shape)
temp_shape[dim] = int(n_chunks * chunk_size)
def noise_sampler(sigma, sigma_next) -> torch.Tensor:
nonlocal select_dim
result = x.new_zeros(temp_shape)
for idx in range(0, dim_size, chunk_size):
throw_exception_if_processing_interrupted()
select_dim[dim] = slice(idx, idx + chunk_size)
result[select_dim] = ns(sigma, sigma_next)[select_dim]
return scale_noise(result[trim_slice], factor, normalized=normalize)
return noise_sampler
@@ -2094,6 +2230,9 @@ class CustomNoiseParametersNoise(CustomNoiseItemBase):
):
factor = self.factor
normalize = self.get_normalize("normalize", normalized)
print(
f"\n****** NS: normalized={normalized}, normalize={normalize}, self.normalize={self.normalize}"
)
orig_shape = x.shape
orig_dtype = x.dtype
orig_device = x.device
@@ -2169,6 +2308,14 @@ class CustomNoiseParametersNoise(CustomNoiseItemBase):
temp_rng_state.set_states()
else:
noise = ns(sigma, sigma_next)
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:
if noise.dtype.is_complex and not orig_dtype.is_complex:
noise = noise.real.mul_(0.5).add_(noise.imag.mul_(0.5))
noise = noise.to(device=orig_device, dtype=orig_dtype)
if fix_invalid:
noise_temp = noise.nan_to_num(0, posinf=0, neginf=0)
noise = noise.nan_to_num_(
@@ -2176,17 +2323,70 @@ class CustomNoiseParametersNoise(CustomNoiseItemBase):
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 NestedNoise(CustomNoiseItemBase):
def clone_key(self, k):
if k == "custom_noises":
return tuple(n.clone() for n in self.custom_noises)
return super().clone_key(k)
def make_noise_sampler(
self,
x,
sigma_min,
sigma_max,
*args,
**kwargs,
):
x_numel = x[0].numel()
x_shape = x.shape
shapes = self.nested_shapes
if x_numel != self.nested_numel:
errstr = f"Sampling latent with shape {x.shape} has {x_numel} element(s) per batch item which does not match the initial reference with shape(s) {shapes} and element(s) {self.nested_numel}"
raise ValueError(errstr)
factor = self.factor
n_shapes = len(shapes)
# batch_size = x.shape[0]
noise_samplers = tuple(
n.make_noise_sampler(
x_part,
*args,
sigma_min=sigma_min,
sigma_max=sigma_max,
**kwargs,
)
for n, x_part in zip(
self.custom_noises[:n_shapes],
comfy_utils.unpack_latents(x, shapes),
strict=True,
)
)
del x
n_samplers = len(noise_samplers)
if not n_samplers:
raise RuntimeError("Internal error: No noise samplers available")
def noise_sampler(*args, **kwargs) -> torch.Tensor:
noise_chunks = (
noise_samplers[nidx % n_samplers](*args, **kwargs)
for nidx in range(n_shapes)
)
if factor != 1:
noise_chunks = (nc.mul_(factor) for nc in noise_chunks)
result = comfy_utils.pack_latents(tuple(noise_chunks))[0]
return (
result
if result.shape[1:] == x_shape[1:]
else result.reshape(-1, *x_shape[1:])
)
return noise_sampler
class BlehOpsNoise(CustomNoiseItemBase):
def __init__(
self,
File diff suppressed because it is too large Load Diff
+55
View File
@@ -0,0 +1,55 @@
from .automata_noise_generator import AutomataNoiseGenerator
from .base import MixedNoiseGenerator, NoiseError, NoiseType
from .collatz_noise_generator import CollatzNoiseGenerator
from .distro_noise_generator import DistroNoiseGenerator
from .novelty_filtered_noise import NoveltyFilteredNoiseGenerator
from .scatternet_filtered_noise_generator import ScatternetFilteredNoiseGenerator
from .simple_noise_generators import (
BrownianNoiseGenerator,
GaussianNoiseGenerator,
GreenTestNoiseGenerator,
HighresPyramidNoiseGenerator,
LaplacianNoiseGenerator,
OneFNoiseGenerator,
PerlinOldNoiseGenerator,
PinkOldNoiseGenerator,
PowerLawNoiseGenerator,
PowerOldNoiseGenerator,
PyramidNoiseGenerator,
PyramidOldNoiseGenerator,
StudentTNoiseGenerator,
UniformNoiseGenerator,
)
from .simulation_noise_generator import SimulationNoiseGenerator
from .voronoi_noise_generator import VoronoiNoiseGenerator
from .wavelet_filtered_noise_generator import WaveletFilteredNoiseGenerator
from .wavelet_noise_generator import WaveletNoiseGenerator
__all__ = (
"AutomataNoiseGenerator",
"BrownianNoiseGenerator",
"CollatzNoiseGenerator",
"DistroNoiseGenerator",
"GaussianNoiseGenerator",
"GreenTestNoiseGenerator",
"HighresPyramidNoiseGenerator",
"LaplacianNoiseGenerator",
"MixedNoiseGenerator",
"NoiseError",
"NoiseType",
"NoveltyFilteredNoiseGenerator",
"OneFNoiseGenerator",
"PerlinOldNoiseGenerator",
"PinkOldNoiseGenerator",
"PowerLawNoiseGenerator",
"PowerOldNoiseGenerator",
"PyramidNoiseGenerator",
"PyramidOldNoiseGenerator",
"ScatternetFilteredNoiseGenerator",
"SimulationNoiseGenerator",
"StudentTNoiseGenerator",
"UniformNoiseGenerator",
"VoronoiNoiseGenerator",
"WaveletFilteredNoiseGenerator",
"WaveletNoiseGenerator",
)
@@ -0,0 +1,325 @@
from __future__ import annotations
import math
from functools import partial
from typing import TYPE_CHECKING, Any
import torch
from comfy.model_management import throw_exception_if_processing_interrupted
from tqdm import trange
from .base import NoiseGenerator
if TYPE_CHECKING:
from collections.abc import Callable
F = torch.nn.functional
# Analytic extension of the Collatz conjecture for floating point numbers.
def continuous_collatz(x: torch.Tensor) -> torch.Tensor:
# f(x) = 1/4 * (2 + 7x - (2 + 5x)*cos(pi*x))
cos_term = (x * torch.pi).cos_()
return x.mul(7).add_(2).sub_(x.mul(5).add_(2).mul_(cos_term)).mul_(0.25)
def generate_spatial_collatz_noise(
batch: int,
channels: int,
height: int,
width: int,
depth: int = None, # Optional 3D depth
steps: int = 20,
num_seeds: int = 15,
device: str = "cpu",
):
is_3d = depth is not None
# 1. Initialize grid
if is_3d:
grid = torch.zeros((batch, channels, depth, height, width), device=device)
norm_dims = [2, 3, 4]
else:
grid = torch.zeros((batch, channels, height, width), device=device)
norm_dims = [2, 3]
# 2. Plant float/negative "seeds"
for b in range(batch):
for c in range(channels):
seed_y = torch.randint(0, height, (num_seeds,))
seed_x = torch.randint(0, width, (num_seeds,))
# Using random floats from -1000 to 1000
seed_vals = (
torch.rand((num_seeds,), dtype=torch.float32, device=device) * 2000.0
) - 1000.0
if is_3d:
seed_z = torch.randint(0, depth, (num_seeds,))
grid[b, c, seed_z, seed_y, seed_x] = seed_vals
else:
grid[b, c, seed_y, seed_x] = seed_vals
# 3. Create spatial diffusion kernel
if is_3d:
# Create a 3x3x3 blurring kernel using outer products
k1d = torch.tensor([1.0, 2.0, 1.0], device=device)
kernel = (k1d.view(3, 1, 1) * k1d.view(1, 3, 1) * k1d.view(1, 1, 3)) / 64.0
kernel = kernel.view(1, 1, 3, 3, 3).repeat(channels, 1, 1, 1, 1)
conv_fn = F.conv3d
else:
# Create a 3x3 blurring kernel
k1d = torch.tensor([1.0, 2.0, 1.0], device=device)
kernel = (k1d.view(3, 1) * k1d.view(1, 3)) / 16.0
kernel = kernel.view(1, 1, 3, 3).repeat(channels, 1, 1, 1)
conv_fn = F.conv2d
# 4. Evolve the grid
for _ in trange(steps, desc="Automata", miniter=25):
# A. Spatial diffusion (spread values into neighboring dimensions)
grid = conv_fn(grid, kernel, padding=1, groups=channels)
# B. Apply Collatz activation
grid = continuous_collatz(grid)
# C. Reset rule: inject new seeds if elements get trapped in low magnitude cycles
trapped_mask = grid.abs() <= 1.5
if trapped_mask.any():
new_seeds = (torch.rand_like(grid) * 200.0) - 100.0
grid = torch.where(trapped_mask, new_seeds, grid)
# D. Internal Instance Normalization to tame the math
mean = grid.mean(dim=norm_dims, keepdim=True)
std = grid.std(dim=norm_dims, keepdim=True) + 1e-5
grid = (grid - mean) / std
# 5. Final Output Normalization
mean = grid.mean(dim=norm_dims, keepdim=True)
std = grid.std(dim=norm_dims, keepdim=True) + 1e-5
return (grid - mean) / std
class AutomataNoiseGenerator(NoiseGenerator):
name = "automata"
blend_function: Callable | None = None
@classmethod
def ng_params(cls):
return super().ng_params() | {
# Evolution mode
# collatz, logistic, sawtooth, lenia, roll
"evolution_mode": "collatz",
# blur, laplacian, crystal
"spread_mode": "blur",
"steps": 20,
"spread_substeps": 3,
"num_seeds": 10,
"depth": 10,
"trapped_threshold": 1.5,
"trapped_interval": 1,
# Controls behavior for trapped elements.
# new - new seed, reset - original seed, mean - replace with mean
"trapped_mode": "reset",
"range_negative": -100.0,
"range_positive": 100.0,
# Absolute value.
"seed_minimum": 1.5,
"noise_sampler_factory": None,
}
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
if not (self.height and self.width):
raise ValueError("Unsupported shape")
self.grid = self.grid_orig = None
self.noise_chunk = None
self.current_depth = 0
def create_grid(self) -> None:
batch, channels = self.batch, self.channels
height, width, depth = self.height, self.width, self.depth
num_seeds = self.num_seeds
is_3d = depth > 0
device, dtype = self.gen_device, self.dtype
total_seeds = batch * channels * num_seeds
# Create flat arrays of coordinates for every single seed
batch_idx = (
torch.arange(batch, device=device)
.view(-1, 1, 1)
.expand(batch, channels, num_seeds)
.flatten()
)
chan_idx = (
torch.arange(channels, device=device)
.view(1, -1, 1)
.expand(batch, channels, num_seeds)
.flatten()
)
y_idx = torch.randint(
0,
height,
(total_seeds,),
device=device,
generator=self.generator,
)
x_idx = torch.randint(
0,
width,
(total_seeds,),
device=device,
generator=self.generator,
)
if is_3d:
grid = torch.zeros(
(batch, channels, depth, height, width),
device=device,
dtype=dtype,
)
else:
grid = torch.zeros(
(batch, channels, height, width),
device=device,
dtype=dtype,
)
seed_vals = (
torch.rand(
(total_seeds,), dtype=dtype, device=device, generator=self.generator
)
* 2000.0
) - 1000.0
seed_vals = seed_vals.abs().clamp_min(self.seed_minimum).copysign(seed_vals)
if is_3d:
z_idx = torch.randint(0, depth, (total_seeds,), device=device)
grid[batch_idx, chan_idx, z_idx, y_idx, x_idx] = seed_vals
else:
grid[batch_idx, chan_idx, y_idx, x_idx] = seed_vals
self.grid = grid
self.initial_grid = grid.clone()
def evolve_step(self, grid: torch.Tensor) -> torch.Tensor:
depth = 0 if grid.ndim < 5 else grid.shape[-3]
# k1d = torch.tensor([1.0, 2.0, 1.0], device=device)
k1d = torch.tensor(
[0.1, 1.0, 0.1],
device=grid.device,
dtype=grid.dtype,
)
if depth > 0:
# Create a 3x3x3 blurring kernel using outer products
kernel = k1d.view(3, 1, 1) * k1d.view(1, 3, 1) * k1d.view(1, 1, 3)
kernel /= kernel.sum()
kernel = kernel.view(1, 1, 3, 3, 3).repeat(self.channels, 1, 1, 1, 1)
else:
# Create a 3x3 blurring kernel
kernel = k1d.view(3, 1) * k1d.view(1, 3)
kernel /= kernel.sum()
kernel = kernel.view(1, 1, 3, 3).repeat(self.channels, 1, 1, 1)
op = partial(
F.conv3d if depth > 0 else F.conv2d,
weight=kernel,
padding=1,
groups=self.channels,
)
# op = partial(F.max_pool3d if depth > 0 else F.max_pool2d, kernel_size=3, stride=1, padding=1)
for _ in range(self.spread_substeps):
grid = grid.lerp(op(grid), 1.0)
# grid = F.max_pool3d(grid, 3, stride=1, padding=1)
# grid = conv_fn(grid, kernel, padding=1, groups=self.channels)
grid = continuous_collatz(grid)
return grid
# return continuous_collatz(grid)
def handle_trapped(
self,
*,
grid: torch.Tensor,
orig_grid: torch.Tensor,
grid_prev: torch.Tensor | None = None,
) -> torch.Tensor:
if self.trapped_threshold == 0:
return grid
mask = grid.abs() < self.trapped_threshold
if grid_prev is not None:
mask &= grid_prev.abs() >= self.trapped_threshold
if not torch.any(mask):
return grid
new_seeds = (torch.rand_like(grid) * 2000.0) - 1000.0
return torch.where(mask, new_seeds, grid)
# return torch.where(mask, orig_grid, grid) if torch.any(mask) else grid
def handle_norm(self, grid: torch.Tensor) -> torch.Tensor:
return grid.clamp(-10000.0, 10000.0)
dims = tuple(range(2, grid.ndim))
gn = grid.clone()
gn /= gn.std(dim=dims, keepdim=True).clamp_min_(1e-06)
return grid.lerp(gn, grid.abs().div_(10000.0).clamp_max_(1.0))
# mask = grid.abs() > 100.0
# return torch.where(mask, grid.lerp(gn, 0.5), grid)
# gn = grid - grid.mean(dim=dims, keepdim=True)
def handle_norm_(self, grid: torch.Tensor) -> torch.Tensor:
# return (grid.abs() % 1000000.0).copysign_(grid)
# mask = grid.abs() > 40000.0
# return torch.where(
# mask,
# (grid.cos() * 1000.0).abs().clamp_min(1.5).copysign(grid),
# grid,
# )
# return torch.where(mask, (grid.abs() % 2000.0).copysign(grid), grid)
# new_seeds = (torch.rand_like(grid) * 2000.0) - 1000.0
# return torch.where(mask, new_seeds, grid)
# return grid * (~mask).to(grid)
mask = grid.abs() > 100000000.0
grid = (grid.abs() % 100000000.0).copysign(grid)
return grid
dims = tuple(range(2, grid.ndim))
std = grid.std(dim=dims, keepdim=True)
std = std.abs().clamp_min_(1e-08).copysign(std)
grid_adj = grid / std
grid_adj -= grid_adj.mean(dim=dims, keepdim=True)
grid = torch.where(mask, grid_adj, grid)
return grid
def evolve(self):
if self.grid is None:
self.create_grid()
grid = self.grid
for i in trange(self.steps, miniters=10, desc="Automata step"):
if i > 1 and (i % 5) == 0:
throw_exception_if_processing_interrupted()
grid_prev = grid
grid = self.evolve_step(grid)
grid = self.handle_trapped(
grid=grid,
orig_grid=self.initial_grid,
grid_prev=grid_prev,
)
grid = self.handle_norm(grid)
self.grid = grid
def reset_grid(self):
self.grid = self.initial_grid = None
self.current_depth = 0
def generate(self, *args) -> torch.Tensor:
if self.grid is None:
self.create_grid()
self.current_depth = 0
self.evolve()
if self.grid.ndim < 5:
return self.grid.clone()
grid = self.grid
self.reset_grid()
return grid
result = self.grid[:, :, self.current_depth, ...].clone()
self.current_depth += 1
if self.current_depth >= self.grid.shape[-3]:
self.reset_grid()
return result
+237
View File
@@ -0,0 +1,237 @@
from __future__ import annotations
from enum import Enum, auto
import torch
from ..utils import (
fallback,
scale_noise,
tensor_to,
)
# ruff: noqa: ANN002, ANN003
class NoiseType(Enum):
BROWNIAN = auto()
COLLATZ = auto()
DISTRO = auto()
GAUSSIAN = auto()
GREEN_TEST = auto()
GREY = auto()
HIGHRES_PYRAMID = auto()
HIGHRES_PYRAMID_AREA = auto()
HIGHRES_PYRAMID_BISLERP = auto()
LAPLACIAN = auto()
ONEF_GREENISH = auto()
ONEF_GREENISH_MIX = auto()
ONEF_PINKISH = auto()
ONEF_PINKISH_MIX = auto()
ONEF_PINKISHGREENISH = auto()
PERLIN = auto()
PINK_OLD = auto()
POWER_OLD = auto()
PYRAMID = auto()
PYRAMID_AREA = auto()
PYRAMID_BISLERP = auto()
PYRAMID_DISCOUNT5 = auto()
PYRAMID_MIX = auto()
PYRAMID_MIX_AREA = auto()
PYRAMID_MIX_BISLERP = auto()
PYRAMID_OLD = auto()
PYRAMID_OLD_AREA = auto()
PYRAMID_OLD_BISLERP = auto()
RAINBOW_INTENSE = auto()
RAINBOW_MILD = auto()
STUDENTT = auto()
UNIFORM = auto()
VELVET = auto()
VIOLET = auto()
VORONOI_FUZZ = auto()
VORONOI_MIX = auto()
WAVELET = auto()
WHITE = auto()
@classmethod
def get_names(cls, default=GAUSSIAN, skip=None):
if default is not None:
if isinstance(default, int):
default = cls(default)
yield default.name.lower()
for nt in cls:
if nt == default or (skip and nt in skip):
continue
yield nt.name.lower()
class NoiseError(Exception):
pass
class NoiseGenerator:
name = "unknown"
MIN_DIMS = 1
MAX_DIMS = 0
def __init__(
self,
x,
**kwargs,
):
if x.ndim < self.MIN_DIMS:
errstr = f"Noise generator {self.name} requires at least {self.MIN_DIMS} dimension(s) but got input with shape {x.shape}"
raise ValueError(errstr)
if self.MAX_DIMS > 0 and x.ndim > self.MAX_DIMS:
errstr = f"Noise generator {self.name} requires at most {self.MAX_DIMS} dimension(s) but got input with shape {x.shape}"
raise ValueError(errstr)
params = self.ng_params()
kwarg_params = params | kwargs
for k in params:
setattr(self, k, kwarg_params.pop(k))
self.options = kwarg_params
self.update_x(x)
@classmethod
def ng_params(cls):
return {
"normalized": True,
"force_normalize": None,
"normalize_dims": None,
"cpu": True,
"generator": None,
}
def update_x(self, x):
self.shape = x.shape
self.batch = self.channels = self.frames = self.height = self.width = None
if x.ndim >= 2:
self.batch, self.channels = x.shape[:2]
if x.ndim > 2:
self.width = x.shape[-1]
if x.ndim > 3:
self.height = x.shape[-2]
if x.ndim == 5:
self.frames = x.shape[-3]
self.device = x.device
self.gen_device = torch.device("cpu") if self.cpu else self.device
self.layout = x.layout
self.dtype = x.dtype
def rand_like(
self,
*,
fun=torch.randn,
cpu=None,
to_device=True,
shape=None,
dtype=None,
layout=None,
device=None,
generator=None,
):
cpu = fallback(cpu, self.cpu)
noise = fun(
*fallback(shape, self.shape),
generator=fallback(generator, self.generator),
dtype=fallback(dtype, self.dtype),
layout=fallback(layout, self.layout),
device=fallback(device, "cpu" if cpu else self.gen_device),
)
if to_device and noise.device != self.device:
noise = tensor_to(noise, self.device)
return noise
def output_hook(self, noise):
if noise.device != self.device:
noise = tensor_to(noise, self.device)
return scale_noise(
noise,
normalized=self.normalized
and (self.force_normalize is None or self.force_normalize is True),
normalize_dims=self.normalize_dims,
)
def pre_hook(self):
pass
def generate(self):
raise NotImplementedError
def __call__(self, *args, **kwargs):
self.pre_hook()
return self.output_hook(self.generate(*args, **kwargs))
def __str__(self):
pretty_params = ", ".join(f"{k}={getattr(self, k)!s}" for k in self.ng_params())
return f"<NoiseGenerator({self.name}): device={self.device}, shape={self.shape}, dtype={self.dtype}, {pretty_params}>"
class FramesToChannelsNoiseGenerator(NoiseGenerator):
MIN_DIMS = 4
MAX_DIMS = 5
def get_adjusted_shape(self):
if self.frames:
return (self.batch, self.channels * self.frames, self.height, self.width)
return (self.batch, self.channels, self.height, self.width)
def fix_output_frames(self, noise):
if not self.frames:
return noise
return noise.reshape(
self.batch,
self.channels,
self.frames,
self.height,
self.width,
)
def rand_like(self, *args, shape=None, **kwargs):
noise = super().rand_like(*args, shape=shape, **kwargs)
if shape is not None:
return noise
adjusted_shape = self.get_adjusted_shape()
if noise.shape != adjusted_shape:
return noise.reshape(*adjusted_shape)
return noise
class MixedNoiseGenerator(NoiseGenerator):
@classmethod
def ng_params(cls):
return super().ng_params() | {
"name": "mixed_noise",
"normalized": True,
"pass_args": frozenset(("cpu",)),
"noise_mix": (),
"output_fun": None,
}
def __init__(self, x, *args, **kwargs):
min_dim = max_dim = None
self.name = kwargs["name"]
for item in kwargs["noise_mix"]:
ng_class = item[0] if isinstance(item, (tuple, list)) else item
cmin, cmax = ng_class.MIN_DIMS, ng_class.MAX_DIMS
min_dim = max(min_dim if min_dim is not None else cmin, cmin)
max_dim = min(max_dim if max_dim is not None else cmax, cmax)
self.MIN_DIMS = min_dim
self.MAX_DIMS = max_dim
super().__init__(x, *args, **kwargs)
ng_list = []
for ng_class, ng_class_kwargs, transform_fun in self.noise_mix:
ng_kwargs = {k: v for k, v in kwargs.items() if k in self.pass_args}
ng_list.append((ng_class(x, **ng_class_kwargs, **ng_kwargs), transform_fun))
self.ng_list = ng_list
def generate(self, *args):
noise = None
for ng, transform_fun in self.ng_list:
new_noise = ng(*args)
if transform_fun is not None:
new_noise = transform_fun(new_noise)
noise = new_noise if noise is None else noise.add_(new_noise)
if self.output_fun is not None:
noise = self.output_fun(noise)
return noise
@@ -0,0 +1,304 @@
# ruff: noqa: ANN002
from __future__ import annotations
import math
from typing import TYPE_CHECKING, ClassVar
import torch
from comfy.model_management import throw_exception_if_processing_interrupted
from .. import utils
from ..utils import fallback, normalize_to_scale, tensor_to
from .base import NoiseGenerator
if TYPE_CHECKING:
from collections.abc import Sequence
F = torch.nn.functional
class CollatzNoiseGenerator(NoiseGenerator):
name = "collatz"
@classmethod
def ng_params(cls):
return super().ng_params() | {
"adjust_scale": False,
"iteration_sign_flipping": True,
"chain_length": (1, 1, 2, 2, 3, 3),
"iterations": 10,
"rmin": -8000.0,
"rmax": 8000.0,
"flatten": False,
"dims": (-1, -1, -2, -2),
# values, ratios, mults, adds
# seed_x_ratios, seed_x_mults, seed_x_adds
# noise_x_ratios, noise_x_mults, noise_x_adds
"output_mode": "values",
"quantile": 0.5,
"quantile_strategy": "clamp",
"noise_dtype": torch.float32,
"integer_math": True,
"even_multiplier": 0.5,
"even_addition": 0.0,
"odd_multiplier": 3.0,
"odd_addition": 1.0,
"add_preserves_sign": True,
"chain_offset": 5,
"break_loops": True,
"seed_mode": "default",
"seed_noise_sampler": None,
"mix_noise_sampler": None,
}
@staticmethod
def _get_iter_slices(n_dims, dim, idx, stride) -> tuple:
result = [slice(None)] * n_dims
result[dim] = slice(idx, None, stride)
return tuple(result)
def _generate_iteration(
self,
*args,
dim: int,
chain_length: int,
flatten: False,
shape=None,
):
dtype, device = self.dtype, self.device
out_shape = shape = fallback(shape, self.shape)
if dim >= len(shape):
raise ValueError("Requested dimension out of range")
rmin, rmax = self.rmin, self.rmax
emul, eadd = self.even_multiplier, self.even_addition
omul, oadd = self.odd_multiplier, self.odd_addition
keepsign = self.add_preserves_sign
intmode = self.integer_math
rmaxsubmin = rmax - rmin
if flatten:
shape = torch.Size((*shape[:dim], math.prod(shape[dim:])))
size = shape[dim]
chain_length = min(size, chain_length)
n_chunks = math.ceil(size / chain_length)
chain_length += self.chain_offset
result_shape = list(shape)
chunk_shape = result_shape.copy()
result_shape[dim] = chain_length * n_chunks
chunk_shape[dim] = n_chunks
result = torch.zeros(result_shape, dtype=self.noise_dtype, device=device)
adds, muls = result.clone(), result.clone()
if self.seed_noise_sampler is not None:
orig_noise = self.seed_noise_sampler(*args)[
tuple(slice(None, sz) for sz in chunk_shape)
].to(result)
if flatten:
orig_noise = orig_noise.flatten(start_dim=dim)
orig_noise = normalize_to_scale(
orig_noise[tuple(slice(None, sz) for sz in chunk_shape)],
1e-06,
1.0,
dim=tuple(range(1, len(chunk_shape))),
)
else:
orig_noise = self.rand_like(
fun=torch.rand,
shape=chunk_shape,
dtype=result.dtype,
)
noise = orig_noise * (rmaxsubmin + 1) + rmin
# Derp.
noise = torch.where(noise == 0, noise.max() / noise.numel(), noise)
if self.seed_mode != "default":
noise = torch.where(
(noise % 2.0) < 1
if self.seed_mode == "force_odd"
else (noise % 2.0) >= 1,
noise + 1,
noise,
)
if noise.device != self.device:
noise = tensor_to(noise, self.device)
slice_0 = self._get_iter_slices(result.ndim, dim, 0, chain_length)
for chainidx in range(chain_length):
if chainidx == 0:
muls[slice_0] = 1.0
result[slice_0] = noise
continue
slice_curr = self._get_iter_slices(result.ndim, dim, chainidx, chain_length)
slice_prev = self._get_iter_slices(
result.ndim,
dim,
chainidx - 1,
chain_length,
)
prev = result[slice_prev]
prev_trunc = utils.trunc_decimals(prev, 2)
need_reset = (
((prev_trunc >= 1.0) & (prev_trunc < 1.001))
| (prev_trunc.abs() < 0.001)
if self.break_loops
else False
)
prev_evens = prev % 2 < 1.0
prev_adds, prev_muls = adds[slice_prev], muls[slice_prev]
muls_next = (
torch.where(
prev_evens,
prev_muls if emul == 1 else prev_muls * emul,
prev_muls if omul == 1 else prev_muls * omul,
)
if emul != 1 or omul != 1
else prev_muls
)
muls[slice_curr] = (
torch.where(need_reset, 1.0, muls_next)
if need_reset is not False
else muls_next
)
curr_muls = muls[slice_curr]
prev_adds_scaled = prev_adds * curr_muls
prev_sign = prev.sign() if keepsign else 1.0
adds_next = (
torch.where(
prev_evens,
prev_adds_scaled
if eadd == 0
else prev_adds_scaled + eadd * prev_sign,
prev_adds_scaled
if oadd == 0
else prev_adds_scaled + oadd * prev_sign,
)
if eadd != 0 or oadd != 0
else prev_adds_scaled
)
adds[slice_curr] = (
torch.where(need_reset, 0.0, adds_next)
if need_reset is not False
else adds_next
)
curr_adds = adds[slice_curr]
result_next = utils.maybe_apply(
(noise * curr_muls).add_(curr_adds),
intmode,
torch.trunc,
)
result[slice_curr] = (
torch.where(need_reset, noise, result_next)
if need_reset is not False
else result_next
)
output_slice = tuple(
slice(None, sz) for sz in (shape if flatten else out_shape)
)
return self._iteration_output(
*args,
result_chains=result,
orig_noise=orig_noise,
noise=noise,
raw_adds=adds,
muls=muls,
chain_length=chain_length,
dim=dim,
output_shape=out_shape,
output_slice=output_slice,
dtype=dtype,
)
def _trim_chain_offset(
self,
t: torch.Tensor,
dim: int,
chain_length: int,
) -> torch.Tensor:
co = self.chain_offset
if co < 1:
return t
chunks = t.split(chain_length, dim)
slices = tuple(
slice(None) if i != dim else slice(co, None) for i in range(t.ndim)
)
return torch.cat(
tuple(chunk[slices] for chunk in chunks),
dim=dim,
)
def _iteration_output(
self,
*args,
result_chains: torch.Tensor,
orig_noise: torch.Tensor,
noise: torch.Tensor,
raw_adds: torch.Tensor,
muls: torch.Tensor,
chain_length: int,
dim: int,
output_shape: Sequence,
output_slice: Sequence,
dtype: str | torch.dtype,
) -> torch.Tensor:
omode = self.output_mode
quantile = self.quantile
noise_exp = noise.repeat_interleave(chain_length, dim)
nadds = raw_adds.div_(noise_exp)
ratios = result_chains / noise_exp
if omode in {"values", "ratios", "seed_x_ratios", "noise_x_ratios"}:
out1 = ratios
elif omode in {"mults", "seed_x_mults", "noise_x_mults"}:
out1 = muls
elif omode in {"adds", "seed_x_adds", "noise_x_adds"}:
out1 = nadds
else:
raise ValueError("Bad output mode")
out1 = self._trim_chain_offset(out1, dim=dim, chain_length=chain_length)
if quantile not in {0, 1}:
out1 = utils.quantile_normalize(
out1,
quantile=quantile,
dim=0,
strategy=self.quantile_strategy,
)
out1 = out1[output_slice].reshape(output_shape).to(dtype=dtype)
if omode in {"ratios", "mults", "adds"}:
return out1
if omode in {"values", "seed_x_ratios", "seed_x_mults", "seed_x_adds"}:
out2 = orig_noise.repeat_interleave(chain_length - self.chain_offset, dim)
elif omode in {"noise_x_ratios", "noise_x_mults", "noise_x_adds"}:
out2 = (
self.rand_like(dtype=out1.dtype)
if self.mix_noise_sampler is None
else self.mix_noise_sampler(*args)
)
out2 = out2[output_slice].reshape(output_shape).to(dtype=dtype)
return out2 * out1
def generate(self, *args):
out_dims = len(self.shape)
dims = tuple(dim if dim >= 0 else out_dims + dim for dim in self.dims)
n_dims, n_chainlens = len(dims), len(self.chain_length)
if not all(0 <= d < out_dims for d in dims):
raise ValueError("Dimension out of range")
dtype, device = self.dtype, self.device
result = torch.zeros(self.shape, dtype=dtype, device=device)
it_scale = 1.0 / self.iterations
for iteration in range(self.iterations):
if iteration > 0 and (iteration % 25) == 0:
# It's soooo slow!
throw_exception_if_processing_interrupted()
temp = self._generate_iteration(
*args,
dim=dims[iteration % n_dims],
chain_length=self.chain_length[iteration % n_chainlens],
flatten=self.flatten,
).mul_(
it_scale
* (-1 if self.iteration_sign_flipping and (iteration & 1) == 1 else 1),
)
result += temp
if self.adjust_scale:
result = normalize_to_scale(
result,
-1.0,
1.0,
dim=tuple(range(1 if result.ndim < 4 else 2, result.ndim)),
)
return result
@@ -0,0 +1,461 @@
# ruff: noqa: ANN002, ANN003
from __future__ import annotations
import torch
from ..utils import quantile_normalize
from .base import NoiseGenerator
class DistroNoiseGenerator(NoiseGenerator):
name = "distro"
simple_distros = frozenset((
"cauchy",
"exponential",
"geometric",
"log_normal",
"normal",
))
def __init__(self, x, *args, **kwargs):
super().__init__(x, *args, **kwargs)
if self.distro not in self.distro_params():
raise ValueError("Bad distro")
_distro_params = None
@classmethod
def distro_params(cls):
if cls._distro_params is not None:
return cls._distro_params
td = torch.distributions
tt = torch.Tensor
cls._distro_params = {
# Simple
"exponential": (
tt.exponential_,
{
"lambd": {
"default": 1.0,
},
},
),
"cauchy": (
tt.cauchy_,
{
"median": {
"default": "0.0",
},
"sigma": {
"default": 1.0,
"min": 0.0,
},
},
),
"geometric": (
tt.geometric_,
{
"p": {
"default": 0.25,
},
},
),
"log_normal": (
tt.log_normal_,
{
"mean": {
"default": 1.0,
},
"std": {
"default": 2.0,
},
},
),
"normal": (
tt.normal_,
{
"mean": {
"default": 0.0,
},
"std": {
"default": 1.0,
},
},
),
# Complex distros
"beta": (
td.Beta,
{
"concentration0": {
"default": "0.5",
},
"concentration1": {
"default": "0.5",
},
},
),
"continuous_bernoulli": (
td.ContinuousBernoulli,
{
"probs": {
"default": "0.5",
},
},
),
"dirichlet": (
td.Dirichlet,
{
"concentration": {
"default": "0.5 0.5",
},
},
),
"fisher_snedecor": (
td.FisherSnedecor,
{
"df1": {
"default": "1.0",
},
"df2": {
"default": "2.0",
},
},
),
"gamma": (
td.Gamma,
{
"concentration": {
"default": "1.0",
},
"rate": {
"default": "1.0",
},
},
),
"gumbel": (
td.Gumbel,
{
"loc": {
"default": "1.0",
},
"scale": {
"default": "2.0",
},
},
),
"inverse_gamma": (
td.InverseGamma,
{
"concentration": {
"default": "1.0",
},
"rate": {
"default": "1.0",
},
},
),
"kumaraswamy": (
td.Kumaraswamy,
{
"concentration0": {
"default": "1.0",
},
"concentration1": {
"default": "1.0",
},
},
),
"laplacian": (
td.Laplace,
{
"loc": {
"default": "0.0",
},
"scale": {
"default": "1.0",
},
},
),
"lkjcholesky": (
td.LKJCholesky,
{
"dim": {
"_ty": "INT",
"default": 3,
},
"concentration": {
"default": "1.0",
},
},
),
"lrmvariate_normal": (
lambda loc, cov_factor, cov_diag: td.LowRankMultivariateNormal(
loc=loc,
cov_factor=cov_factor.reshape(loc.numel(), -1),
cov_diag=cov_diag,
),
{
"loc": {
"default": "0.0 0.0",
},
"cov_factor": {
"default": "1.0 0.0",
},
"cov_diag": {
"default": "1.0 1.0",
},
},
),
"mvariate_normal": (
lambda loc, cov_multiplier=1.0: td.MultivariateNormal(
loc=loc,
covariance_matrix=torch.eye(
loc.numel(),
dtype=loc.dtype,
device=loc.device,
).mul_(cov_multiplier),
),
{
"loc": {
"default": "0.0 0.0",
},
"cov_multiplier": {
"default": 1.0,
},
},
),
"pareto": (
td.Pareto,
{
"scale": {
"default": "1.0",
},
"alpha": {
"default": "1.0",
},
},
),
"poisson": (
td.Poisson,
{
"rate": {
"default": "1.5",
},
},
),
"relaxed_bernoulli": (
td.RelaxedBernoulli,
{
"temperature": {
"default": 0.75,
},
"probs": {
"default": "0.66",
},
},
),
"relaxed_onehotcategorical": (
td.RelaxedOneHotCategorical,
{
"temperature": {
"default": 1.5,
},
"probs": {
"default": "0.33 0.66",
},
},
),
"studentt": (
td.StudentT,
{
"loc": {
"default": "0.0",
},
"scale": {
"default": "1.0",
},
"df": {
"default": "1.0",
},
},
),
"uniform": (
td.Uniform,
{
"low": {
"default": 0.0,
},
"high": {
"default": 1.0,
},
},
),
"vonmises": (
td.VonMises,
{
"loc": {
"default": "1.0",
},
"concentration": {
"default": "1.0",
},
},
),
"weibull": (
td.Weibull,
{
"scale": {
"default": "1.0",
},
"concentration": {
"default": "1.0",
},
},
),
"wishart": (
lambda df, cov_size=2, cov_multiplier=1.0: td.Wishart(
df=df,
covariance_matrix=torch.eye(
int(cov_size),
dtype=df.dtype,
device=df.device,
).mul_(cov_multiplier),
),
{
"df": {
"default": "2.0",
},
"cov_size": {
"_ty": "INT",
"default": 2,
},
"cov_multiplier": {
"default": 1.0,
},
},
),
}
return cls._distro_params
_build_params = None
@classmethod
def build_params(cls):
if cls._build_params is not None:
return cls._build_params
cls._build_params = {
f"{tykey}_{pkey}": pval
for tykey, tyval in cls.distro_params().items()
for pkey, pval in tyval[1].items()
if not pkey.startswith("_")
}
return cls._build_params
_ng_params = None
@classmethod
def ng_params(cls):
if cls._ng_params is not None:
return cls._ng_params
dparams = {
k: v["default"]
for k, v in cls.build_params().items()
if not k.startswith("_")
}
cls._ng_params = (
super().ng_params()
| {
"distro": "normal",
"quantile_norm": 0.85,
"quantile_norm_flatten": True,
"quantile_norm_dim": 1,
"quantile_norm_pow": 0.5,
"quantile_norm_fac": 1.0,
"result_index": "-1",
}
| dparams
)
return cls._ng_params
def norm_output(self, noise):
if noise.ndim > len(self.shape):
if noise.shape[: len(self.shape)] != self.shape:
errstr = f"Unexpected shape when normalizing distro({self.distro}) noise! Output shape={self.shape}, noise shape={noise.shape}, generator dump: {self}"
raise RuntimeError(errstr)
selfdims = len(self.shape)
result_index = self.result_index
if not isinstance(result_index, (tuple, list)):
result_index = (result_index,)
ri_len = len(result_index)
if ri_len == 0:
raise ValueError("When result_index is a list, it must not be empty")
trim_count = 0
while noise.ndim > selfdims:
idx = result_index[trim_count % ri_len]
if idx < 0:
idx = noise.shape[-1] + idx
noise = noise[..., max(0, min(noise.shape[-1] - 1, idx))]
trim_count += 1
return (
quantile_normalize(
noise,
quantile=self.quantile_norm,
dim=self.quantile_norm_dim,
flatten=self.quantile_norm_flatten,
nq_fac=self.quantile_norm_fac,
pow_fac=self.quantile_norm_pow,
)
.reshape(self.shape)
.contiguous()
)
def distro_param(self, val, *, simple_fun=None):
if isinstance(val, torch.Tensor):
return simple_fun(val) if simple_fun is not None else val
if isinstance(val, str):
val = tuple(float(v) for v in val.split(None))
if simple_fun is not None:
if isinstance(val, (float, int)):
return simple_fun(val)
if len(val) > 1:
raise ValueError("Couldn't return result as float")
return simple_fun(val[0])
if not isinstance(val, (tuple, list)):
val = (val,)
return torch.tensor(
val,
dtype=self.dtype,
device=self.gen_device,
)
def get_distro_kwargs(self, distro, ddef, *, simple=False):
return {
k: self.distro_param(
getattr(self, f"{distro}_{k}"),
simple_fun=None
if not simple and k != "dim"
else (int if k == "dim" else float),
)
for k in ddef
}
def generate(self, *_args):
distro = self.distro
dfun, ddef = self.distro_params()[distro]
is_simple = distro in self.simple_distros
dkwargs = self.get_distro_kwargs(distro, ddef, simple=is_simple)
if is_simple:
noise = torch.empty(
*self.shape,
device=self.gen_device,
dtype=self.dtype,
layout=self.layout,
)
noise = dfun(noise, **dkwargs)
else:
dobj = dfun(**dkwargs)
noise = (
dobj.rsample if getattr(dobj, "has_rsample", False) else dobj.sample
)(self.shape)
return self.norm_output(noise)
@@ -0,0 +1,186 @@
# ruff: noqa: ANN002, ANN003
from __future__ import annotations
import math
from functools import partial
from typing import TYPE_CHECKING
import torch
from .base import NoiseGenerator
if TYPE_CHECKING:
from collections.abc import Callable
F = torch.nn.functional
def sum_rms_blend(
a: torch.Tensor,
b: torch.Tensor,
t: torch.Tensor | float = 1.0,
*,
orig_shape: torch.Size | tuple[int, ...],
dims_a: tuple[int, ...] = (1,),
dims_b: tuple[int, ...] = (-1, -2),
) -> torch.Tensor:
rms_a = a / math.prod(orig_shape[d] for d in dims_a) ** 0.5
rms_b = b / math.prod(orig_shape[d] for d in dims_b) ** 0.5
variance_a = rms_a.pow_(2.0)
variance_b = rms_b.pow_(2.0)
result = variance_a
result += variance_b * t
result /= 1.0 + t
result **= 0.5
return result
def metrics_blend(
a: torch.Tensor,
b: torch.Tensor,
t: torch.Tensor | float = 1.0,
*,
orig_shape: torch.Size | tuple[int, ...],
dims_a: tuple[int, ...] = (-1, -2),
dims_b: tuple[int, ...] = (1,),
use_rms: bool = True,
rms_power: float = 2.0,
) -> torch.Tensor:
count_a = math.prod(orig_shape[d] for d in dims_a)
count_b = math.prod(orig_shape[d] for d in dims_b)
denom_a = count_a ** (1 / rms_power) if use_rms else count_a
denom_b = count_b ** (1 / rms_power) if use_rms else count_b
curr_a = a / denom_a
curr_b = b / denom_b
if use_rms:
curr_a **= rms_power
curr_b = curr_b.pow_(rms_power) * t
result = curr_b.add_(curr_a)
result /= 1.0 + t
return result.pow_(1.0 / rms_power) if use_rms else result
class NoveltyFilteredNoiseGenerator(NoiseGenerator):
name = "novelty"
initial_noise_state: torch.Tensor | None = None
noise_state: torch.Tensor | None = None
blend_function: Callable | None = None
@classmethod
def ng_params(cls):
return super().ng_params() | {
"skip_initial": 1,
"iters_per_call": 1,
"dim_groups": ((1,), (-1, -2)),
"blend_ratio": 1.0,
"blend_function": None,
"update_blend_ratio": 1.0,
"update_blend_function": None,
"noise_sampler": None,
}
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
if self.blend_function is None:
raise ValueError("Missing blend function!")
def generate(self, *args) -> torch.Tensor:
ng = (
partial(self.noise_sampler, *args) if self.noise_sampler else self.rand_like
)
noise_state = self.noise_state
had_state = self.noise_state is not None
it_counter = 0 if had_state else 0 - self.skip_initial
its_call = max(1, self.iters_per_call)
bf = self.blend_function
blend_ratio = self.blend_ratio
update_blend_ratio = self.update_blend_ratio
ubf = self.update_blend_function
if ubf is None:
# Linear weighted average
def ubf(a: torch.Tensor, b: torch.Tensor, t: float) -> torch.Tensor:
return (b * t).add_(a).div_(1.0 + abs(t))
call_initial_noise = None
curr_noise = None
call_noise_state = None
while it_counter < its_call:
if noise_state is None:
noise_state = ng()
self.initial_noise_state = noise_state.clone()
self.noise_state = noise_state.clone()
continue
curr_noise = ng()
if call_initial_noise is None:
call_initial_noise = curr_noise.clone()
seen = {id(curr_noise)}
for ortho_target in (
self.initial_noise_state,
call_initial_noise if it_counter > 0 else None,
noise_state,
):
tid = id(ortho_target)
if ortho_target is None or tid in seen:
continue
seen.add(tid)
curr_noise = bf(ortho_target, curr_noise, blend_ratio)
curr_noise -= ortho_target
it_counter += 1
if it_counter < 1:
self.noise_state = curr_noise.clone()
noise_state = curr_noise
continue
if call_noise_state is None:
call_noise_state = curr_noise
else:
call_noise_state = ubf(call_noise_state, curr_noise, update_blend_ratio)
if call_noise_state is None:
raise RuntimeError("Unexpected unpopulated call_noise_state!")
# self.noise_state = call_noise_state.clone()
self.noise_state = ubf(noise_state, call_noise_state, update_blend_ratio)
return call_noise_state
# def generate(self, *args) -> torch.Tensor:
# ng = (
# partial(self.noise_sampler, *args) if self.noise_sampler else self.rand_like
# )
# noise_state = self.noise_state
# had_state = self.noise_state is not None
# it_counter = 0 if had_state else 0 - self.skip_initial
# its_call = self.iters_per_call
# bf = self.blend_function
# blend_ratio = self.blend_ratio
# update_blend_ratio = self.update_blend_ratio
# ubf = self.update_blend_function
# if ubf is None or True:
# # Linear weighted average
# def ubf(a: torch.Tensor, b: torch.Tensor, t: float) -> torch.Tensor:
# return (b * t).add_(a).div_(1.0 + abs(t))
# call_initial_noise = None
# while it_counter < its_call:
# if noise_state is None:
# noise_state = ng()
# self.initial_noise_state = noise_state.clone()
# continue
# curr_noise = ng()
# it_counter += 1
# if it_counter < 1:
# noise_state = curr_noise
# continue
# if call_initial_noise is None and self.iters_per_call > 1:
# call_initial_noise = curr_noise.clone()
# for ortho_target in (
# self.initial_noise_state,
# call_initial_noise if it_counter > 0 else None,
# noise_state,
# ):
# if ortho_target is None:
# continue
# curr_noise = bf(ortho_target, curr_noise, blend_ratio)
# curr_noise -= ortho_target
# # curr_noise = bf(noise_state, ng(), blend_ratio).sub_(noise_state)
# noise_state = ubf(noise_state, curr_noise, update_blend_ratio)
# self.noise_state = noise_state.clone()
# return noise_state
@@ -0,0 +1,173 @@
# ruff: noqa: ANN002, ANN003
from __future__ import annotations
import math
import torch
from .. import utils
from ..wavelet_functions import ptwav
from .base import FramesToChannelsNoiseGenerator
F = torch.nn.functional
class ScatternetFilteredNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "scatternetfilter"
MIN_DIMS = 4
MAX_DIMS = 4
def __init__(self, *args, **kwargs):
if ptwav is None:
raise RuntimeError(
"Scatternet noise requires the pytorch_wavelets package to be installed in your Python environment",
)
super().__init__(*args, **kwargs)
if self.output_mode not in {
"channels",
"channels_adjusted",
"channels_scaled",
"flat",
"flat_adjusted",
"flat_scaled",
}:
raise ValueError("Bad output mode")
scatkwargs = {
"mode": self.mode,
"biort": "near_sym_b_bp" if self.use_symmetric_filter else self.biort,
}
if self.scatternet_order == 2:
scatkwargs["qshift"] = (
"qshift_b_bp" if self.use_symmetric_filter else self.qshift
)
self.scatternet = ptwav.ScatLayerj2(**scatkwargs)
elif self.scatternet_order == 1:
self.scatternet = ptwav.ScatLayer(**scatkwargs)
else:
self.scatternet = torch.nn.Sequential(
*(
ptwav.ScatLayer(**scatkwargs)
for _ in range(abs(self.scatternet_order))
),
)
@classmethod
def ng_params(cls):
return super().ng_params() | {
"mode": "symmetric",
"magbias": 1e-02,
"use_symmetric_filter": False,
"biort": "near_sym_a",
"qshift": "qshift_a",
"output_offset": 0.0,
"scatternet_order": 1,
"per_channel_scatternet": False,
"output_mode": "channels_adjusted",
# If None, uses probselect when available, otherwise bilinear.
"upscale_mode": None,
"noise_sampler": None,
}
def _fix_shape(self, noise, adjusted_shape):
if self.frames:
noise = noise.reshape(
self.batch,
self.channels * self.frames,
self.height,
self.width,
)
elif noise.shape != adjusted_shape:
noise = noise.reshape(*adjusted_shape)
return noise
def generate(self, *args):
adjusted_shape = self.get_adjusted_shape()
scaled = self.output_mode.endswith("_scaled")
adjusted = scaled or self.output_mode.endswith("_adjusted")
order = abs(self.scatternet_order)
order_spatial_compensation = 2**order
output_mode = (
self.output_mode.split("_", 1)[0] if adjusted else self.output_mode
)
spatial_compensation = 1 if adjusted else order_spatial_compensation
if self.noise_sampler is None:
temp_shape = (
(
*adjusted_shape[:2],
adjusted_shape[-2] * spatial_compensation,
adjusted_shape[-1] * spatial_compensation,
)
if spatial_compensation != 1
else adjusted_shape
)
noise = self.rand_like(shape=temp_shape)
else:
noise = self.noise_sampler(*args)
if scaled:
upscale_mode = self.upscale_mode
if upscale_mode is None:
upscale_mode = (
"probselect"
if "probselect" in utils.UPSCALE_METHODS
else "bilinear"
)
noise = utils.scale_samples(
noise,
adjusted_shape[-1] * order_spatial_compensation,
adjusted_shape[-2] * order_spatial_compensation,
mode=upscale_mode,
)
if self.scatternet_order == 0:
return self.fix_output_frames(noise)
self.scatternet = self.scatternet.to(device=self.device, dtype=self.dtype)
if self.per_channel_scatternet:
# To C, B, 1, H, W
noise = torch.stack(
tuple(
self.scatternet(noise[:, chan : chan + 1])
for chan in range(self.channels)
),
dim=0,
)
else:
# To 1, B, C, H, W
noise = self.scatternet(noise)[None]
base_channels = 1 if self.per_channel_scatternet else self.channels
if output_mode == "flat":
noise = noise.reshape(noise.shape[0], self.batch, -1)
initial_size = math.prod(
self.shape[(2 if self.per_channel_scatternet else 1) :],
)
elif adjusted:
initial_size = base_channels
else:
initial_size = base_channels * ((2**order) ** 2)
increment = 1 if output_mode == "flat" else base_channels
out_size = noise.shape[2]
offset_size = (out_size - initial_size) / increment
output_offset = self.output_offset
if output_offset == 0 or abs(output_offset) >= 1:
output_offset = int(output_offset)
if output_offset < 0:
output_offset = (offset_size + 1) + output_offset
else:
if output_offset < 0:
output_offset += 1.0
output_offset = round(offset_size * output_offset)
base_idx = int(output_offset * increment)
# print(
# f"\nSCAT: shape={noise.shape}, adj_shape={adjusted_shape}, offset={output_offset}, initial_size={initial_size}, out_size={out_size}, offset_size={offset_size}, incr={increment}, base_idx={base_idx}",
# )
noise = noise[:, :, base_idx : base_idx + initial_size]
# print(f"\nSCAT2: {noise.shape}")
noise = (
noise.squeeze(2).movedim(0, 1) if self.per_channel_scatternet else noise[0]
)
# print(f"\nSCAT3: {noise.shape}")
if output_mode == "channels":
noise = noise[..., : self.height, : self.width]
# print(
# f"\nSCAT4: {noise.shape} -> {adjusted_shape} -- numel: {noise.numel()}, adjnumel={math.prod(adjusted_shape)}",
# )
return noise.reshape(adjusted_shape).contiguous()
@@ -0,0 +1,698 @@
from __future__ import annotations
import math
import operator
from typing import Callable
import torch
from comfy.k_diffusion import sampling
from torch import FloatTensor, Generator, Tensor
from torch.distributions import Laplace, StudentT
from .. import utils
from ..utils import safe_pow, tensor_to
# ruff: noqa: D413, D417, D212, ANN002, ANN003
from .base import FramesToChannelsNoiseGenerator, NoiseError, NoiseGenerator
class GaussianNoiseGenerator(NoiseGenerator):
name = "gaussian"
@classmethod
def ng_params(cls):
return super().ng_params() | {"normalized": False}
def generate(self, *_args):
return self.rand_like()
class BrownianNoiseGenerator(NoiseGenerator):
name = "brownian"
def __init__(self, x, *args, **kwargs):
super().__init__(x, *args, **kwargs)
seed = self.options.get("seed")
sigma_min = self.options.get("sigma_min")
sigma_max = self.options.get("sigma_max")
if sigma_min is None or sigma_max is None:
raise ValueError("Brownian noise requires sigma_min and sigma_max")
self.brownian_tree_ns = sampling.BrownianTreeNoiseSampler(
x,
sigma_min,
sigma_max,
seed=seed,
cpu=self.cpu,
)
@classmethod
def ng_params(cls):
return super().ng_params() | {"normalized": False}
def generate(self, *args):
return self.brownian_tree_ns(*args)
class PerlinOldNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "perlin_old"
@classmethod
def ng_params(cls):
return super().ng_params() | {
"div_fac": 2.0,
"iterations": 2,
"blend_mode": "lerp",
}
@staticmethod
def get_positions(block_shape: tuple[int, int]) -> Tensor:
"""
Generate position tensor.
Arguments:
block_shape -- (height, width) of position tensor
Returns:
position vector shaped (1, height, width, 1, 1, 2)
"""
bh, bw = block_shape
return torch.stack(
torch.meshgrid(
[(torch.arange(b) + 0.5) / b for b in (bw, bh)],
indexing="xy",
),
-1,
).view(1, bh, bw, 1, 1, 2)
@staticmethod
def unfold_grid(vectors: Tensor) -> Tensor:
"""
Unfold vector grid to batched vectors.
Arguments:
vectors -- grid vectors
Returns:
batched grid vectors
"""
batch_size, _channels, gpy, gpx = vectors.shape
return (
torch.nn.functional.unfold(vectors, (2, 2))
.view(batch_size, 2, 4, -1)
.permute(0, 2, 3, 1)
.view(batch_size, 4, gpy - 1, gpx - 1, 2)
)
@staticmethod
def smooth_step(t: Tensor) -> Tensor:
"""
Smooth step function [0, 1] -> [0, 1].
Arguments:
t -- input values (any shape)
Returns:
output values (same shape as input values)
"""
return t * t * (3.0 - 2.0 * t)
@classmethod
def perlin_noise_tensor(
cls,
vectors: Tensor,
positions: Tensor,
step: Callable | None = None,
blend=torch.lerp,
) -> Tensor:
"""
Generate perlin noise from batched vectors and positions.
Arguments:
vectors -- batched grid vectors shaped (batch_size, 4, grid_height, grid_width, 2)
positions -- batched grid positions shaped (batch_size or 1, block_height, block_width, grid_height or 1, grid_width or 1, 2)
Keyword Arguments:
step -- smooth step function [0, 1] -> [0, 1] (default: `smooth_step`)
Raises:
NoiseError: if position and vector shapes do not match
Returns:
(batch_size, block_height * grid_height, block_width * grid_width)
"""
if step is None:
step = cls.smooth_step
batch_size = vectors.shape[0]
# grid height, grid width
gh, gw = vectors.shape[2:4]
# block height, block width
bh, bw = positions.shape[1:3]
for i in range(2):
if positions.shape[i + 3] not in {1, vectors.shape[i + 2]}:
msg = f"Blocks shapes do not match: vectors ({vectors.shape[1]}, {vectors.shape[2]}), positions {gh}, {gw})"
raise NoiseError(msg)
if positions.shape[0] not in {1, batch_size}:
msg = f"Batch sizes do not match: vectors ({vectors.shape[0]}), positions ({positions.shape[0]})"
raise NoiseError(msg)
vectors = vectors.view(batch_size, 4, 1, gh * gw, 2)
positions = positions.view(positions.shape[0], bh * bw, -1, 2)
step_x = step(positions[..., 0])
step_y = step(positions[..., 1])
row0 = blend(
(vectors[:, 0] * positions).sum(dim=-1),
(vectors[:, 1] * (positions - positions.new_tensor((1, 0)))).sum(dim=-1),
step_x,
)
row1 = blend(
(vectors[:, 2] * (positions - positions.new_tensor((0, 1)))).sum(dim=-1),
(vectors[:, 3] * (positions - positions.new_tensor((1, 1)))).sum(dim=-1),
step_x,
)
noise = blend(row0, row1, step_y)
return (
noise.view(
batch_size,
bh,
bw,
gh,
gw,
)
.permute(0, 3, 1, 4, 2)
.reshape(batch_size, gh * bh, gw * bw)
)
@classmethod
def perlin_noise(
cls,
grid_shape: tuple[int, int],
out_shape: tuple[int, int],
batch_size: int = 1,
blend=torch.lerp,
generator: Generator | None = None,
*args,
**kwargs,
) -> Tensor:
"""
Generate perlin noise with given shape. `*args` and `**kwargs` are forwarded to `Tensor` creation.
Arguments:
grid_shape -- Shape of grid (height, width).
out_shape -- Shape of output noise image (height, width).
Keyword Arguments:
batch_size -- (default: {1})
generator -- random generator used for grid vectors (default: {None})
Raises:
NoiseError: if grid and out shapes do not match
Returns:
Noise image shaped (batch_size, height, width)
"""
# grid height and width
gh, gw = grid_shape
# output height and width
oh, ow = out_shape
# block height and width
bh, bw = oh // gh, ow // gw
if oh != bh * gh:
msg = f"Output height {oh} must be divisible by grid height {gh}"
raise NoiseError(msg)
if ow != bw * gw != 0:
msg = f"Output width {ow} must be divisible by grid width {gw}"
raise NoiseError(msg)
angle = torch.empty(
[batch_size] + [s + 1 for s in grid_shape],
*args,
**kwargs,
).uniform_(to=2.0 * math.pi, generator=generator)
# random vectors on grid points
vectors = cls.unfold_grid(
torch.stack((torch.cos(angle), torch.sin(angle)), dim=1),
)
# positions inside grid cells [0, 1)
positions = tensor_to(cls.get_positions((bh, bw)), vectors)
return cls.perlin_noise_tensor(vectors, positions, blend=blend).squeeze(0)
def generate(self, *_args):
blend = utils.BLENDING_MODES[self.blend_mode]
noise = self.rand_like(fun=torch.rand).div_(self.div_fac)
channels, height, width = noise.shape[1:]
for _ in range(self.iterations):
noise += self.perlin_noise(
(height, self.width),
(height, width),
batch_size=channels,
blend=blend,
dtype=noise.dtype,
layout=noise.layout,
device=noise.device,
)
return self.fix_output_frames(noise)
class UniformNoiseGenerator(NoiseGenerator):
name = "uniform"
@classmethod
def ng_params(cls):
return super().ng_params() | {
"normalized": False,
"sub_fac": 0.5,
"mul_fac": 3.46,
"mean_fac": 0.0,
}
def generate(self, *_args):
return (
self.rand_like(fun=torch.rand)
.sub_(self.sub_fac)
.mul_(self.mul_fac)
.add_(self.mean_fac)
)
class HighresPyramidNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "highres_pyramid"
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
if self.noise_generator is None:
self.noise_generator = UniformNoiseGenerator(
*args,
**(kwargs | {"normalized": self.normalize_noise}),
)
@classmethod
def ng_params(cls):
return super().ng_params() | {
"normalized": True,
"discount": 0.7,
"upscale_mode": "bilinear",
"iterations": 4,
"noise_generator": None,
"normalize_noise": False,
}
def generate(self, s, sn):
adjusted_shape = self.get_adjusted_shape()
b, c, h, w = adjusted_shape
orig_w, orig_h = w, h
noise = self.noise_generator(s, sn).reshape(*adjusted_shape)
rs = (
torch.rand(
self.iterations,
dtype=torch.float32,
generator=self.generator,
).cpu()
* 2
+ 2
)
for i in range(self.iterations):
r = rs[i].item()
h, w = min(orig_h * 15, int(h * (r**i))), min(orig_w * 15, int(w * (r**i)))
noise += utils.scale_samples(
tensor_to(torch.randn(b, c, h, w, generator=self.generator), noise),
orig_w,
orig_h,
mode=self.upscale_mode,
).mul_(self.discount**i)
if h >= orig_h * 15 or w >= orig_w * 15:
break # Lowest resolution is 1x1
return self.fix_output_frames(noise)
class PyramidOldNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "pyramid_old"
@classmethod
def ng_params(cls):
return super().ng_params() | {
"discount": 0.8,
"iterations": 5,
"upscale_mode": "nearest-exact",
"normalized": False,
}
def generate(self, *_args):
adjusted_shape = self.get_adjusted_shape()
b, c, h, w = adjusted_shape
orig_h, orig_w = h, w
noise = torch.zeros(
size=adjusted_shape,
dtype=self.dtype,
layout=self.layout,
device=self.gen_device,
)
r = 1
for i in range(self.iterations):
r *= 2
noise += utils.scale_samples(
torch.normal(
mean=0,
std=0.5**i,
size=(b, c, h * r, w * r),
dtype=noise.dtype,
layout=noise.layout,
generator=self.generator,
device=noise.device,
),
orig_w,
orig_h,
mode=self.upscale_mode,
).mul_(self.discount**i)
return self.fix_output_frames(noise)
class PyramidNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "pyramid"
@classmethod
def ng_params(cls):
return super().ng_params() | {
"discount": 0.7,
"upscale_mode": "bilinear",
"iterations": 10,
"iteration_offset": 0,
"iteration_step": 1,
"reverse_scale": False,
"reverse_size_h": False,
"reverse_size_w": False,
"base_h": 2.0,
"multiplier_h": 2.0,
"base_w": 2.0,
"multiplier_w": 2.0,
"legacy_r": False,
"size_min": 1,
"size_max_pct": 2.0,
"include_size_limit": False,
"high_res_mode": False,
}
# Original implementatino modified from https://wandb.ai/johnowhitaker/multires_noise/reports/Multi-Resolution-Noise-for-Diffusion-Model-Training--VmlldzozNjYyOTU2
def generate(self, *_args):
noise = self.rand_like()
b, c, h, w = noise.shape
orig_w, orig_h = w, h
size_min = max(1, self.size_min)
max_h = max(1, int(orig_h * self.size_max_pct))
max_w = max(1, int(orig_w * self.size_max_pct))
eps = 1e-02
op = operator.mul if self.high_res_mode else operator.truediv
if self.legacy_r:
def get_r(_i: int) -> float:
return torch.rand(1, generator=self.generator).cpu().item()
else:
rs = torch.rand(self.iterations, generator=self.generator).cpu().tolist()
def get_r(i: int) -> float:
return rs[i]
for i in range(
self.iteration_offset,
self.iterations + self.iteration_offset,
self.iteration_step,
):
rev_i = self.iterations - i - 1
r = get_r(i)
rh = r * self.multiplier_h + self.base_h
rw = r * self.multiplier_w + self.base_w
ih = rev_i if self.reverse_size_h else i
iw = rev_i if self.reverse_size_w else i
h = max(1, min(max_h, int(op(h, max(eps, rh**ih)))))
w = max(1, min(max_w, int(op(w, max(eps, rw**iw)))))
size_limit = h <= size_min or w <= size_min or h >= max_h or w >= max_w
if not self.include_size_limit and size_limit:
break
scale = self.discount ** (rev_i if self.reverse_scale else i)
if scale == 0:
continue
noise += utils.scale_samples(
torch.randn(
b,
c,
h,
w,
device=noise.device,
layout=noise.layout,
dtype=noise.dtype,
),
orig_w,
orig_h,
mode=self.upscale_mode,
).mul_(scale)
if size_limit:
break
return self.fix_output_frames(noise)
class StudentTNoiseGenerator(NoiseGenerator):
name = "studentt"
@classmethod
def ng_params(cls):
return super().ng_params() | {
"loc": 0,
"scale": 0.2,
"df": 1,
"quantile_fac": 0.75,
"pow_fac": 0.5,
"nq_fac": 1.0,
"normalized": False,
}
def generate(self, *_args):
noise = StudentT(loc=self.loc, scale=self.scale, df=self.df).rsample(self.shape)
nq = torch.quantile(
noise.flatten(start_dim=1).abs(),
self.quantile_fac,
dim=-1,
)
nq_shape = tuple(nq.shape) + (1,) * (noise.ndim - nq.ndim)
nq = nq.mul_(self.nq_fac).reshape(*nq_shape)
noise = noise.clamp_(-nq, nq)
return noise.abs().pow_(self.pow_fac).copysign_(noise)
class GreenTestNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "green_test"
MIN_DIMS = 4
MAX_DIMS = 5
@classmethod
def ng_params(cls):
return super().ng_params() | {
"scale_fac": 1.0,
"x_pow": 2.0,
"y_pow": 2.0,
"x_multiplier": 1.0,
"y_multiplier": 1.0,
"power_base": 1.0,
"inv_power": 0.5,
"restore_sign_power": False,
"restore_sign_x": False,
"restore_sign_y": False,
}
def generate(self, *_args):
noise = self.rand_like()
scale = self.scale_fac / max(1, self.width * self.height)
fy, fx = (
torch.fft.fftfreq(sz, device=noise.device, dtype=noise.dtype)
for sz in (self.height, self.width)
)
fx = safe_pow(fx, self.x_pow, restore_sign=self.restore_sign_x, in_place=True)
fy = safe_pow(fy, self.y_pow, restore_sign=self.restore_sign_y, in_place=True)
if self.x_multiplier != 1:
fx *= self.x_multiplier
if self.y_multiplier != 1:
fy *= self.y_multiplier
power = fy[:, None] + fx
inv_power = self.inv_power * self.inv_power
power = safe_pow(
power,
inv_power,
restore_sign=self.restore_sign_power,
in_place=True,
)
coord_0 = self.power_base**self.inv_power
if coord_0 == 0 or not math.isfinite(coord_0):
coord_0 = 1.0
power = power.masked_fill_((power == 0) | (~power.isfinite()), coord_0)
power[0, 0] = coord_0
noise *= scale
noise = torch.fft.ifft2(torch.fft.fft2(noise).div_(power))
return self.fix_output_frames(noise.real)
class PinkOldNoiseGenerator(NoiseGenerator):
name = "pink_old"
@classmethod
def ng_params(cls):
return super().ng_params() | {"alpha": 2.0, "k": 1.0, "freq": 1.0}
# Completely wrong implementation here.
def generate(self, *_args):
spectral_density = self.k / self.freq**self.alpha
return self.rand_like() * spectral_density
def frequency_scaled_noise(
x: torch.Tensor,
*,
x_is_noise: bool = False,
base_power: float = 0.5,
alpha: float,
) -> torch.Tensor:
h, w = x.shape[-2:]
fh = torch.fft.fftfreq(h, device=x.device).unsqueeze(-1)
fw = torch.fft.fftfreq(w, device=x.device).unsqueeze(0)
p = (fh**2 + fw**2).pow_(base_power * alpha)
p[0, 0] = 1.0**alpha
noise = x if x_is_noise else torch.randn_like(x)
noise_fft = torch.fft.fftn(noise, dim=(-2, -1))
p = p.to(noise_fft.dtype).expand(*((1,) * (x.ndim - 2)), h, w)
noise_fft /= p
noise_fft[..., 0, 0] = 0.0
noise = torch.fft.ifftn(noise_fft, dim=(-2, -1)).real.to(x.dtype)
noise /= noise.std(dim=tuple(range(1, x.ndim)), keepdim=True).clamp_min_(1e-06)
return noise
class OneFNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "onef"
MIN_DIMS = 4
MAX_DIMS = 5
@classmethod
def ng_params(cls):
return super().ng_params() | {
"alpha": 2.0,
"k": 1.0,
"hfac": 1.0,
"wfac": 1.0,
"base_power": 1.0,
"use_sqrt": True,
# None or or float, alternative to use_sqrt with custom power.
"power": None,
"x_pow": 2.0,
"y_pow": 2.0,
}
# Original implementation referenced from: https://github.com/WASasquatch/PowerNoiseSuite
def generate(self, *_args):
noise = self.rand_like()
freq_x, freq_y = (
torch.fft.fftfreq(sz, fac, device=noise.device, dtype=noise.dtype)
for sz, fac in ((self.height, self.hfac), (self.width, self.wfac))
)
freq_x **= self.x_pow
freq_y **= self.y_pow
fx, fy = torch.meshgrid(freq_x, freq_y, indexing="ij")
power = fx + fy
power **= self.alpha / -2.0
if self.k not in {0, 1}:
power *= 1 / self.k
noise_fft = torch.fft.fftn(noise)
user_power = 0.5 if self.use_sqrt else self.power
if isinstance(user_power, float):
power **= user_power
coord_0 = self.base_power
if coord_0 == 0 or not math.isfinite(coord_0):
coord_0 = 1.0
power = power.masked_fill_((power == 0) | (~power.isfinite()), coord_0)
power = (
power.to(dtype=noise_fft.dtype)
.unsqueeze(0)
.expand(self.batch, 1, self.height, self.width)
)
noise_fft /= power
noise = torch.fft.ifftn(noise_fft).real
return self.fix_output_frames(noise)
class PowerLawNoiseGenerator(NoiseGenerator):
name = "powerlaw"
@classmethod
def ng_params(cls):
return super().ng_params() | {
"alpha": 2.0,
"div_max_dims": None,
"use_sign": False,
"use_div_max_abs": True,
}
# Referenced from: https://github.com/WASasquatch/PowerNoiseSuite
def generate(self, *_args):
noise = self.rand_like()
modulation = torch.abs(noise) ** self.alpha
noise = (torch.sign(noise) if self.use_sign else noise).mul_(modulation)
if self.div_max_dims is not None:
noise /= torch.amax(
torch.abs(noise) if self.use_div_max_abs else noise,
keepdim=True,
dim=self.div_max_dims,
)
return noise
class LaplacianNoiseGenerator(NoiseGenerator):
name = "laplacian"
@classmethod
def ng_params(cls):
return super().ng_params() | {"loc": 0, "scale": 1.0, "div_fac": 4.0}
def generate(self, *_args):
noise = self.rand_like().div_(self.div_fac)
noise += tensor_to(
Laplace(loc=self.loc, scale=self.scale).rsample(self.shape),
noise.device,
)
return noise
class PowerOldNoiseGenerator(NoiseGenerator):
name = "power_old"
@classmethod
def ng_params(cls):
return super().ng_params() | {"alpha": 2, "k": 1, "normalized": False}
def generate(self, *_args):
tensor = self.rand_like()
fft = torch.fft.fft2(tensor)
freq = torch.arange(
1,
len(fft) + 1,
dtype=tensor.dtype,
layout=tensor.layout,
device=tensor.device,
).reshape(
(len(fft),) + (1,) * (tensor.dim() - 1),
)
spectral_density = self.k / freq**self.alpha
noise = torch.rand(
tensor.shape,
device=tensor.device,
layout=tensor.layout,
dtype=tensor.dtype,
).mul_(spectral_density)
mean = torch.mean(noise, dim=(-2, -1), keepdim=True)
std = torch.std(noise, dim=(-2, -1), keepdim=True)
return noise.sub_(mean).div_(std)
@@ -0,0 +1,597 @@
# ruff: noqa: ANN002, ANN003
from __future__ import annotations
import itertools
import math
from typing import Any
import torch
from tqdm import tqdm
from .base import NoiseGenerator
F = torch.nn.functional
class SimulationNoiseGenerator(NoiseGenerator):
name = "simulation"
MIN_DIMS = 4
MAX_DIMS = 4
@classmethod
def ng_params(cls, *, no_super: bool = False):
result = {
# multi_octave, power_law, band_pass
"spectral_mode": "multi_octave",
# curl, projection, basis
"field_mode": "basis",
"depth_mode": "reset",
"channel_mode": "stacked",
"band_shape": "log_gaussian",
"dims": (),
"base_k": 0.0,
"power_law_beta": 1.0,
"depth": 64,
"initial_depth": 0,
"max_depth": -1,
# reset, wrap, bounce
"octaves": 5,
"lacunarity": 2.0,
"gain": 0.5,
# log_gaussian, raised_cosine
"log_gaussian_sigma": 0.3,
# (float, float, float)
"band_pass_low": 0.00001,
"band_pass_high": 1.0,
"anisotropy": (),
"normalized": False,
"noise_sampler_factory_h": None,
"noise_sampler_factory_w": None,
"noise_sampler_factory_z": None,
}
return result if no_super else super().ng_params() | result
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.noise_chunk = None
cm = self.channel_mode
self.depth_increment = 1
if cm in {"over_depth", "over_depth_alt"}:
self.depth_increment = math.ceil(self.channels / 3)
elif cm.startswith("over_depth_"):
self.depth_increment = self.channels
else:
self.depth_increment = 1
if self.initial_depth < 0:
self.initial_depth = self.depth + self.initial_depth
if self.initial_depth < 0:
raise ValueError("Initial depth out of range")
self.initial_depth = min(self.depth - 1, self.initial_depth)
if self.max_depth < 0:
self.max_depth = self.depth + self.max_depth
if self.max_depth < 0:
raise ValueError("Max depth out of range")
self.max_depth = min(self.depth - 1, self.max_depth)
self.current_depth = self.initial_depth
self.direction = 1
self.cdtype = (
(torch.complex128 if self.dtype == torch.float64 else torch.complex64)
if not self.dtype.is_complex
else self.dtype
)
self.eff_batch = (
self.batch
if cm not in {"stacked", "flat"}
else self.batch * math.ceil(self.channels / 3)
)
ns_shape = torch.Size(
(
self.eff_batch,
self.depth * self.depth_increment,
self.height,
self.width,
)
)
def gaussian_noise_sampler(*_args: Any) -> torch.Tensor:
return torch.randn(ns_shape, dtype=self.cdtype, device=self.gen_device).to(
device=self.device,
)
self.noise_samplers = tuple(
factory.make_noise_sampler(
torch.zeros(ns_shape, device=self.gen_device, dtype=self.cdtype),
cpu=self.cpu,
normalized=False,
)
if factory is not None
else gaussian_noise_sampler
for factory in (
self.noise_sampler_factory_z,
self.noise_sampler_factory_h,
self.noise_sampler_factory_w,
)
)
def _k_grids(self, *, shape: tuple, dims: tuple = (-3, -2, -1)) -> tuple:
"""Creates k-space grids."""
return torch.meshgrid(
*(
torch.fft.fftfreq(
shape[dim],
d=1.0,
device=self.device,
dtype=self.dtype if not self.dtype.is_complex else torch.float64,
).to(dtype=self.dtype)
for dim in dims
),
indexing="ij",
)
@staticmethod
def _radial_k(*ks: torch.Tensor) -> torch.Tensor:
"""Calculates the radial distance in k-space."""
return sum(kt**2 for kt in ks).sqrt_()
@staticmethod
def _raised_cosine_band(
k: torch.Tensor,
k_lo: float,
k_hi: float,
) -> torch.Tensor:
"""A raised cosine spectral band filter."""
kc = 0.5 * (k_lo + k_hi)
hw = 0.5 * (k_hi - k_lo) + 1e-12
t = (k - kc) / hw
return torch.where(
t.abs() <= 1.0,
0.5 * (1.0 + torch.cos(math.pi * t)),
torch.zeros_like(k),
)
def _log_gaussian_band(
self,
k: torch.Tensor,
k_lo: float,
k_hi: float,
) -> torch.Tensor:
"""A log-Gaussian spectral band filter."""
k_center = math.sqrt(k_lo * k_hi)
log_k = torch.log(torch.clamp(k, min=1e-12))
log_center = math.log(k_center)
return torch.exp(-0.5 * ((log_k - log_center) / self.log_gaussian_sigma) ** 2)
def _handle_band_shape(
self,
k_rad: torch.Tensor,
k_low: float,
k_high: float,
) -> torch.Tensor:
if self.band_shape == "raised_cosine":
return self._raised_cosine_band(k_rad, k_low, k_high)
if self.band_shape == "log_gaussian":
return self._log_gaussian_band(k_rad, k_low, k_high)
errstr = f"Bad band shape mode {self.band_shape}"
raise ValueError(errstr)
def _make_wk(
self,
k_rad: torch.Tensor,
sizes: tuple,
*,
eps: float = 1e-09,
) -> torch.Tensor:
def wk_out(wk: torch.Tensor) -> torch.Tensor:
wk[k_rad == 0] = 0.0
return wk
if self.spectral_mode == "power_law":
return wk_out((k_rad + eps).pow_(-self.power_law_beta))
if self.spectral_mode == "band_pass":
if self.band_pass_low >= self.band_pass_high:
raise ValueError(
"band_pass_high must be greater than band_pass_low in band_pass spectral mode.",
)
return wk_out(
self._handle_band_shape(k_rad, self.band_pass_low, self.band_pass_high),
)
if self.octaves == 0:
# Ones where k_rad is non-zero, otherwise zero.
return (k_rad != 0).to(k_rad)
base_k = 2 * math.pi / max(1, min(sizes)) if self.base_k == 0 else self.base_k
wk = torch.zeros_like(k_rad)
for o in range(self.octaves):
k_lo = base_k * (self.lacunarity**o)
k_hi = base_k * (self.lacunarity ** (o + 1))
band = self._handle_band_shape(k_rad, k_lo, k_hi)
wk += (self.gain**o) * band
return wk_out(wk)
def _handle_field_projection(
self,
*,
k_grids_orig: tuple,
wk: torch.Tensor,
ns_args: tuple | list,
**_kwargs,
):
n_dims = len(k_grids_orig)
n_samplers = len(self.noise_samplers)
f_fs = tuple(
self.noise_samplers[ns_idx % n_samplers](*ns_args)
.to(
device=self.device,
)
.mul_(wk)
for ns_idx in range(n_dims)
)
# --- Perform the Helmholtz projection using the UN SCALED grids ---
k_sq_proj = self._radial_k(*k_grids_orig) ** 2
k_dot_f = sum(k_p * f_f for k_p, f_f in zip(k_grids_orig, f_fs))
inv_k_sq = torch.where(k_sq_proj == 0, 0.0, 1.0 / k_sq_proj)
k_grid_scale = k_dot_f.mul_(inv_k_sq)
return tuple(
f_f - k_grid * k_grid_scale for f_f, k_grid in zip(f_fs, k_grids_orig)
)
def _handle_field_curl(
self,
*,
k_rad: torch.Tensor,
k_grids_orig: tuple,
wk: torch.Tensor,
ns_args: tuple | list,
**_kwargs,
):
n_dims = len(k_grids_orig)
nd_fixup = int(self.field_mode != "curl_ndim")
# The potential filter still uses the scaled k_rad for spectral shaping
inv_k_rad = torch.where(k_rad == 0, 0.0, 1.0 / k_rad)
wk_potential = wk * inv_k_rad
# The curl operator (i*k) MUST use the original, un-scaled grids
i_k_grids = tuple(
(1j * k_grid).to(dtype=self.cdtype) for k_grid in k_grids_orig
)
n_samplers = len(self.noise_samplers)
g_fs = tuple(
self.noise_samplers[ns_idx % n_samplers](*ns_args)
.to(
device=self.device,
)
.mul_(wk_potential)
for ns_idx in range(n_dims if n_dims != 2 else 1)
)
# --- Case 1: 2D Curl (Curl of a SCALAR potential) ---
# This is the fundamental building block.
if n_dims * nd_fixup == 2:
# We only need one scalar potential field G.
g_f = g_fs[0]
ikx, iky = i_k_grids
# F = (dG/dy, -dG/dx) -> F_f = (iky*G_f, -ikx*G_f)
return (iky * g_f, -ikx * g_f)
# --- Case 2: 3D Curl (The classic cross-product) ---
# This is a special, unique case.
if n_dims * nd_fixup == 3:
gz_f, gy_f, gx_f = g_fs
ikz, iky, ikx = i_k_grids
# F_f = i*k x G_f
return (
ikx * gy_f - iky * gx_f, # z component
ikz * gx_f - ikx * gz_f, # y component
iky * gz_f - ikz * gy_f, # x component
)
# --- Case 3: N-D Curl (Pragmatic construction) ---
# We build the N-D field by summing 2D curls on orthogonal planes.
# We need N potential fields, but we will use them in pairs.
f_f_outputs = [torch.zeros_like(g_fs[0]) for _ in range(n_dims)]
# Iterate over pairs of dimensions (0,1), (2,3), etc.
for i in range(n_dims // 2):
idx1 = i * 2
idx2 = i * 2 + 1
g1_f = g_fs[idx1]
g2_f = g_fs[idx2]
ik1 = i_k_grids[idx1]
ik2 = i_k_grids[idx2]
# Perform a 2D-like curl on the (G1, G2) plane
# This is a bit abstract, but we are creating rotation in the 1-2 plane.
# f_f_outputs[idx1] = ik2 * g1_f - ik1 * g2_f
# f_f_outputs[idx2] = ik1 * g2_f - ik2 * g1_f
f_f_outputs[idx1] = ik2 * g1_f - ik1 * g2_f
f_f_outputs[idx2] = -ik1 * g1_f - ik2 * g2_f
return tuple(f_f_outputs)
_handle_field_curl_ndim = _handle_field_curl
def _handle_field_basis(
self,
*,
k_grids_orig: tuple,
wk: torch.Tensor,
ns_args: tuple | list,
**_kwargs,
) -> tuple:
n_dims = len(k_grids_orig)
nd_fixup = int(self.field_mode != "basis_ndim")
k_rad_orig = self._radial_k(*k_grids_orig)
# Normalize the original k vector
k_norm_components = tuple(
torch.where(k_rad_orig == 0, 0.0, k / k_rad_orig) for k in k_grids_orig
)
# --- Case 1: 2D (simple and fast) ---
if n_dims * nd_fixup == 2:
# The basis is a single vector perpendicular to k: u = (-ky, kx)
kn_y, kn_x = k_norm_components
basis_vectors = [
(-kn_x, kn_y),
] # A list containing one basis vector (a tuple)
num_random_fields = 1
# --- Case 2: 3D (fast cross-product method) ---
elif n_dims * nd_fixup == 3:
num_random_fields = 2
kn_z, kn_y, kn_x = k_norm_components
ez = torch.tensor([0.0, 0.0, 1.0], device=self.device, dtype=self.dtype)
is_parallel = (kn_x.abs() < 1e-6) & (kn_y.abs() < 1e-6)
ux = torch.where(is_parallel, 0.0, kn_y * ez[2] - kn_z * ez[1])
uy = torch.where(is_parallel, -kn_z, kn_z * ez[0] - kn_x * ez[2])
uz = torch.where(is_parallel, kn_x, kn_x * ez[1] - kn_y * ez[0])
u_mag = torch.sqrt(ux**2 + uy**2 + uz**2)
inv_u_mag = torch.where(u_mag == 0, 0.0, 1.0 / u_mag)
ux, uy, uz = ux * inv_u_mag, uy * inv_u_mag, uz * inv_u_mag
# u = (uz, uy, ux)
u = (ux, uy, uz)
vx = kn_y * u[2] - kn_z * u[1]
vy = kn_z * u[0] - kn_x * u[2]
vz = kn_x * u[1] - kn_y * u[0]
v = (vz, vy, vx)
basis_vectors = [u, v]
# --- Case 3: N-D (General Gram-Schmidt process) ---
else:
num_random_fields = n_dims - 1
basis_vectors = []
# Start with the standard basis vectors (e.g., [1,0,0], [0,1,0], [0,0,1])
for i in range(n_dims):
# Create a standard basis vector e_i
e_i = [torch.zeros_like(k_rad_orig) for _ in range(n_dims)]
e_i[i] = torch.ones_like(k_rad_orig)
# Start with v = e_i and make it orthogonal to k
v = list(e_i)
dot_k = sum(
v_comp * k_comp for v_comp, k_comp in zip(v, k_norm_components)
)
v = [
v_comp - dot_k * k_comp
for v_comp, k_comp in zip(v, k_norm_components)
]
# Make it orthogonal to all previously found basis vectors
for b in basis_vectors:
dot_b = sum(v_comp * b_comp for v_comp, b_comp in zip(v, b))
v = [v_comp - dot_b * b_comp for v_comp, b_comp in zip(v, b)]
# Normalize the new basis vector
v_mag = torch.sqrt(sum(comp**2 for comp in v))
# Only add the vector if it's not a zero vector
if torch.any(v_mag > 1e-6):
inv_v_mag = torch.where(v_mag == 0, 0.0, 1.0 / v_mag)
v = [comp * inv_v_mag for comp in v]
basis_vectors.append(tuple(v))
if len(basis_vectors) == num_random_fields:
break
# --- Field Construction (works for all cases) ---
# Generate N-1 independent random complex scalar fields
n_samplers = len(self.noise_samplers)
random_fields = tuple(
self.noise_samplers[ns_idx % n_samplers](*ns_args)
.to(device=self.device)
.mul_(wk)
for ns_idx in range(num_random_fields)
)
# Initialize the final field components to zero
f_f_outputs = [torch.zeros_like(random_fields[0]) for _ in range(n_dims)]
# Project each random field onto its corresponding basis vector and sum them up
for i in range(num_random_fields):
a_f = random_fields[i]
basis_vec = basis_vectors[i]
for j in range(n_dims):
f_f_outputs[j] += a_f * basis_vec[j]
return tuple(f_f_outputs)
_handle_field_basis_ndim = _handle_field_basis
def calculate_spectral_divergence_3d(
self,
field: torch.Tensor,
*,
debug: bool = False,
) -> torch.Tensor:
if field.ndim != 5:
errstr = f"Field must be 5d, got shape {field.shape}"
raise ValueError(errstr)
C = field.shape[1]
if C != 3:
errstr = f"Field must have 3 channels, but has {C}"
raise ValueError(errstr)
cdtype = torch.complex128 if field.dtype == torch.float64 else torch.complex64
KX, KY, KZ = (t.to(field) for t in self._k_grids(shape=field.shape))
fx_f = torch.fft.fftn(field[:, 0, ...], dim=(-3, -2, -1))
fy_f = torch.fft.fftn(field[:, 1, ...], dim=(-3, -2, -1))
fz_f = torch.fft.fftn(field[:, 2, ...], dim=(-3, -2, -1))
div_f = (
(1j * KX.to(cdtype)) * fx_f
+ (1j * KY.to(cdtype)) * fy_f
+ (1j * KZ.to(cdtype)) * fz_f
)
result = torch.fft.ifftn(div_f, dim=(-3, -2, -1)).real
divergences = result.abs_().mean(dim=tuple(range(1, result.ndim)))
if not debug:
return divergences
prettydivs = ", ".join(
f"{dm:.5f}" for dm in divergences.detach().cpu().tolist()
)
tqdm.write(
f"Simulation noise: Input shape: {field.shape}, Mean Absolute Divergences (per batch): {prettydivs}",
)
return divergences
def generate_field(
self,
batch: int,
height: int,
width: int,
*,
ns_args: tuple | list,
) -> torch.Tensor:
depth = self.depth * self.depth_increment
eff_shape = torch.Size((batch, 3, depth, height, width))
# 1. Create the UN SCALED k-grids for the projection operator.
k_grids_orig = k_grids = self._k_grids(shape=eff_shape)
# 2. Create a separate set of k-grids for spectral shaping.
# These can be scaled by the anisotropy factors.
if self.anisotropy and not all(v in {0, 1} for v in self.anisotropy):
n_anisotropy = len(self.anisotropy)
anisotropy = tuple(
1.0 if idx >= n_anisotropy else self.anisotropy[idx] for idx in range(3)
)
k_grids = tuple(
k_p if a in {None, 0, 1} else k_p / a
for k_p, a in itertools.zip_longest(k_grids, anisotropy)
)
# 3. Calculate radial k for the spectral envelope using the SCALED grids.
k_rad = self._radial_k(*k_grids)
# Build the multi-octave spectral envelope (Wk) using the anisotropic k_rad
wk = self._make_wk(k_rad, sizes=(depth, height, width))
field_handler = getattr(self, f"_handle_field_{self.field_mode}", None)
if field_handler is None:
errstr = f"Bad field mode {self.field_mode}"
raise ValueError(errstr)
f_f_outputs = field_handler(
k_rad=k_rad,
k_grids=k_grids,
k_grids_orig=k_grids_orig,
wk=wk,
ns_args=ns_args,
)
# Inverse FFT to transform the field back to the spatial domain
fields = tuple(
torch.fft.ifftn(f_proj, dim=(-3, -2, -1)).real
for f_proj in reversed(f_f_outputs)
)
field = torch.stack(fields, dim=1)
self.calculate_spectral_divergence_3d(field, debug=True)
rms = torch.sqrt(torch.mean(field**2))
if rms > 1e-9:
field /= rms
return field
def generate(self, *args) -> torch.Tensor:
cm = self.channel_mode
if self.noise_chunk is None:
self.noise_chunk = self.generate_field(
self.eff_batch,
self.height,
self.width,
ns_args=args,
).to(dtype=self.dtype)
self.current_depth = self.initial_depth
depth_from = self.current_depth * self.depth_increment
depth_to = depth_from + self.depth_increment
if cm == "stacked":
noise = self.noise_chunk[:, :, self.current_depth]
noise = torch.cat(
tuple(
noise[bidx * self.batch : bidx * self.batch + self.batch]
for bidx in range(noise.shape[0] // self.batch)
),
dim=2,
)
elif cm == "flat":
noise = self.noise_chunk[:, :, self.current_depth]
noise = noise.flatten()[: math.prod(self.shape)]
elif cm in {"over_depth", "over_depth_alt"}:
noise = self.noise_chunk[:, :, depth_from:depth_to]
if cm == "over_depth":
noise = noise.movedim(2, 1)
elif cm == "over_depth_avg":
noise = self.noise_chunk[:, :, depth_from:depth_to].mean(dim=1)
elif cm.startswith("over_depth_"):
channel_lookup = {"h": 0, "w": 1, "z": 2}
mathop = cm[-5:-2]
if mathop in {"add", "sub", "mul", "div"}:
chan1, chan2 = channel_lookup[cm[-7]], channel_lookup[cm[-1]]
noise1 = self.noise_chunk[:, chan1 : chan1 + 1, depth_from:depth_to]
noise2 = self.noise_chunk[:, chan2 : chan2 + 1, depth_from:depth_to]
if mathop == "sub":
noise = noise1 - noise2
elif mathop == "add":
noise = noise1 + noise2
elif mathop == "mul":
noise = noise1 * noise2
elif mathop == "div":
noise = noise1 / (noise2 + 1e-07)
else:
chan = channel_lookup[cm[-1]]
noise = self.noise_chunk[:, chan : chan + 1, depth_from:depth_to]
else:
raise ValueError("Bad channel mode")
self.current_depth += 1 * self.direction
if self.current_depth > self.max_depth or self.current_depth < 0:
dm = self.depth_mode
if dm == "reset":
self.noise_chunk = None
elif dm == "wrap":
self.current_depth = self.initial_depth
elif dm == "bounce":
if self.depth < 2:
raise ValueError("Bounce depth mode requires depth of at least 2")
self.direction = -self.direction
self.current_depth += 2 * self.direction
return (
noise.reshape(self.batch, -1, self.height, self.width)[
:,
: self.channels,
]
.clone()
.contiguous()
)
@@ -0,0 +1,628 @@
# ruff: noqa: ANN002, ANN003
from __future__ import annotations
from typing import Callable
import torch
from .. import utils
from .base import NoiseGenerator
F = torch.nn.functional
# With help from ChatGPT.
class VoronoiNoiseGenerator(NoiseGenerator):
name = "voronoi"
MIN_DIMS = 4
MAX_DIMS = 4
voronoi_distance_modes = frozenset((
"angle_sigmoid",
"angle_tanh",
"angle",
"chebyshev",
"euclidean",
"fractal_norm",
"fuzz",
"manhatten",
"minkowski",
"quadratic",
"weight",
))
voronoi_result_modes = frozenset((
"cellid",
"diff",
"diff2",
"f",
"f1",
"f2",
"f3",
"f4",
"fractal_norm",
"fuzz",
"inv_f",
"inv_f1",
"inv_f2",
"inv_f3",
"inv_f4",
"gradient_magnitude",
"median_distance",
"ridge",
"softmin",
))
@classmethod
def ng_params(cls, *, no_super: bool = False):
result = {
"n_points": (32,),
"distance_mode": ("euclidean",),
"z_initial": 0.0,
"z_increment": 1.0,
"z_max": 100000,
"z_max_mode": "reset",
# None or numeric
"z_range": None,
"result_mode": ("f1",),
"octaves": 1,
# same_features or new_features
"octave_mode": "same_features",
"lacunarity": 2.0, # scale increase per octave
"gain": 0.5, # amplitude decrease per octave
"initial_amplitude": 1.0,
"initial_scale": 1.0,
"noise_sampler_factory": None,
"normalized": False,
}
return result if no_super else super().ng_params() | result
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.feature_points = self.grid_xyz = None
self.noise_samplers = None
self.n_points = tuple(max(2, val) for val in self.n_points)
def voronoi_reset(self, *args):
self.z_curr = self.z_initial
octave_range = tuple(
range(self.octaves if self.octave_mode == "new_features" else 1),
)
if self.noise_sampler_factory is not None and self.noise_samplers is None:
self.noise_samplers = tuple(
self.noise_sampler_factory.make_noise_sampler(
torch.zeros(
self.batch,
self.channels,
self.n_points[octave % len(self.n_points)],
3,
device=self.gen_device,
dtype=self.dtype,
),
cpu=self.cpu,
normalized=False,
)
for octave in octave_range
)
self.feature_points = tuple(
(
torch.rand(
self.batch,
self.channels,
self.n_points[octave % len(self.n_points)],
3,
device=self.gen_device,
dtype=self.dtype,
)
if self.noise_samplers is None
else utils.normalize_to_scale(
self.noise_samplers[octave](*args),
target_min=0.0,
target_max=1.0,
dim=(-1, -2),
)
).to(device=self.device)
for octave in octave_range
)
if self.grid_xyz is not None:
return
y = torch.linspace(
0,
self.height - 1,
self.height,
device=self.device,
dtype=self.dtype,
)
x = torch.linspace(
0,
self.width - 1,
self.width,
device=self.device,
dtype=self.dtype,
)
self.grid_xyz = torch.stack(
torch.meshgrid(y, x, indexing="ij"),
dim=-1,
) / torch.tensor(
(self.height, self.width),
device=self.device,
)
def get_feature_points(self, octave: int) -> torch.Tensor:
result = self.feature_points[octave % len(self.feature_points)]
odd_octave = (octave % 2) == 1
om = self.octave_mode
if (om == "same_invert_odd" and odd_octave) or (
om == "same_invert_even" and not odd_octave
):
return 1.0 - result
if octave > 0 and om in {"same_roll_chan_up", "same_roll_chan_down"}:
return torch.roll(
result,
(-1 if om == "same_roll_chan_up" else 1) * (octave % 3),
dims=(1,),
)
if octave > 0 and om in {"same_roll_dir_up", "same_roll_dir_down"}:
return torch.roll(
result,
(-1 if om == "same_roll_dir_up" else 1) * (octave % 3),
dims=(3,),
)
return result
def get_distance_mode(self, octave: int) -> torch.Tensor:
return self.distance_mode[octave % len(self.distance_mode)]
def get_result_mode(self, octave: int) -> torch.Tensor:
return self.result_mode[octave % len(self.result_mode)]
def voronoi_call_mode(
self,
name: str,
*,
result: bool,
args: list | tuple = (),
kwargs: dict | None = None,
) -> torch.Tensor:
name = name.strip().lower()
modes = self.voronoi_result_modes if result else self.voronoi_distance_modes
mode_label = "result" if result else "distance"
if name not in modes:
errstr = f"Bad Voronoi {mode_label} mode {name}"
raise ValueError(errstr)
kwargs = (
{}
if kwargs is None
else {
k[1:] if k.startswith("_") and len(k) > 1 else k: v
for k, v in kwargs.items()
}
)
return getattr(self, f"_voronoi_{mode_label}_{name}")(*args, **kwargs)
@staticmethod
def _voronoi_distance_euclidean(d: torch.Tensor, **_kwargs) -> torch.Tensor:
return d.pow(2).sum(dim=-1).sqrt_()
@staticmethod
def _voronoi_distance_manhatten(d: torch.Tensor, **_kwargs) -> torch.Tensor:
return d.pow(2).sum(dim=-1).sqrt_()
@staticmethod
def _voronoi_distance_chebyshev(d: torch.Tensor, **_kwargs) -> torch.Tensor:
return d.abs().amax(dim=-1)
@staticmethod
def _voronoi_distance_minkowski(
d: torch.Tensor,
*,
p: float | str = 3.0,
**_kwargs,
) -> torch.Tensor:
p = float(p)
return d.abs().pow(p).sum(dim=-1).pow(1 / p)
@staticmethod
def _voronoi_distance_quadratic(d: torch.Tensor, **_kwargs) -> torch.Tensor:
return d.pow(2).sum(dim=-1)
@staticmethod
def _voronoi_distance_angle(
d: torch.Tensor,
*,
idx: int | str = 2,
**_kwargs,
) -> torch.Tensor:
return (
torch.nn.functional.normalize(d, dim=-1)[..., int(idx)]
.clamp_(-1.0, 1.0)
.acos_()
)
@staticmethod
def _voronoi_distance_angle_tanh(
d: torch.Tensor,
*,
idx: int | str = 2,
**_kwargs,
) -> torch.Tensor:
return torch.nn.functional.normalize(d, dim=-1)[..., int(idx)].tanh_().acos_()
@staticmethod
def _voronoi_distance_angle_sigmoid(
d: torch.Tensor,
*,
idx: int | str = 2,
**_kwargs,
) -> torch.Tensor:
return (
torch.nn.functional.normalize(d, dim=-1)[..., int(idx)]
.sigmoid_()
.mul_(2)
.sub_(1)
.acos_()
)
def _voronoi_distance_weight(
self,
d: torch.Tensor,
*args,
name: str = "euclidean",
h: float | str = 1.0,
w: float | str = 1.0,
z: float | str = 0.25,
**kwargs,
) -> torch.Tensor:
weights = d.new_tensor((float(h), float(w), float(z)))
return self.voronoi_call_mode(
name,
result=False,
args=(d * weights, *args),
kwargs=kwargs,
)
def _voronoi_distance_fractal_norm(
self,
d: torch.Tensor,
*args,
name: str = "euclidean",
mode: str = "sin",
scale: str | float = 0.1,
multiplier: str | float = 10.0,
**kwargs,
) -> torch.Tensor:
if mode == "sin":
fun = torch.sin
elif mode == "cos":
fun = torch.cos
else:
raise ValueError(
"Bad mode parameter for fractal_norm distance mode, must be one of: sin, cos",
)
adjustment = float(scale) * fun(d * float(multiplier))
return self.voronoi_call_mode(
name,
result=False,
args=(d + adjustment, *args),
kwargs=kwargs,
)
def _voronoi_distance_fuzz(
self,
*args,
name: str = "euclidean",
fuzz: float | str = 0.25,
**kwargs,
) -> torch.Tensor:
fuzz = float(fuzz)
result = self.voronoi_call_mode(name, result=False, args=args, kwargs=kwargs)
rmin, rmax = result.aminmax()
fuzz = max(abs(rmin.item()), abs(rmax.item())) * fuzz
result += (
torch.rand(result.shape, device=self.gen_device, dtype=result.dtype)
.mul_(fuzz * 2)
.sub_(fuzz)
.to(device=result.device)
)
return utils.normalize_to_scale(result, rmin.item(), rmax.item(), dim=(-2, -1))
@staticmethod
def _voronoi_result_f(
_d: torch.Tensor,
*,
get_sorted: Callable,
idx: int | str = 0,
**_kwargs,
) -> torch.Tensor:
return get_sorted()[..., int(idx)]
def _voronoi_result_f1(self, *args, **kwargs) -> torch.Tensor:
return self._voronoi_result_f(*args, idx=0, **kwargs)
def _voronoi_result_f2(self, *args, **kwargs) -> torch.Tensor:
return self._voronoi_result_f(*args, idx=1, **kwargs)
def _voronoi_result_f3(self, *args, **kwargs) -> torch.Tensor:
return self._voronoi_result_f(*args, idx=2, **kwargs)
def _voronoi_result_f4(self, *args, **kwargs) -> torch.Tensor:
return self._voronoi_result_f(*args, idx=3, **kwargs)
def _voronoi_result_inv_f(self, *args, eps=1e-06, **kwargs) -> torch.Tensor:
return 1.0 / (self._voronoi_result_f(*args, **kwargs) + eps)
def _voronoi_result_inv_f1(self, *args, **kwargs) -> torch.Tensor:
return self._voronoi_result_inv_f(*args, idx=0, **kwargs)
def _voronoi_result_inv_f2(self, *args, **kwargs) -> torch.Tensor:
return self._voronoi_result_inv_f(*args, idx=1, **kwargs)
def _voronoi_result_inv_f3(self, *args, **kwargs) -> torch.Tensor:
return self._voronoi_result_inv_f(*args, idx=2, **kwargs)
def _voronoi_result_inv_f4(self, *args, **kwargs) -> torch.Tensor:
return self._voronoi_result_inv_f(*args, idx=3, **kwargs)
def _voronoi_result_diff(
self,
*args,
idx1: int | str = 0,
idx2: int | str = 1,
**kwargs,
) -> torch.Tensor:
val1, val2 = (
self._voronoi_result_f(*args, idx=i, **kwargs) for i in (idx1, idx2)
)
return val2 - val1
def _voronoi_result_diff2(
self,
*args,
idx1: int | str = 0,
idx2: int | str = 1,
**kwargs,
) -> torch.Tensor:
val1, val2 = (
self._voronoi_result_f(*args, idx=i, **kwargs) for i in (idx1, idx2)
)
return (val2 - val1) / (val2 + val1 + 1e-06)
@staticmethod
def _voronoi_result_cellid(d, *_args, **_kwargs) -> torch.Tensor:
cellids = d.argmin(dim=-1).to(dtype=d.dtype)
return (cellids / cellids.max()).add_(1.0)
def _voronoi_result_ridge(
self,
*args,
name: str = "diff",
exp: float | str = -10.0,
**kwargs,
) -> torch.Tensor:
return 1.0 - (
float(exp)
* self.voronoi_call_mode(name, result=True, args=args, kwargs=kwargs)
)
@staticmethod
def _voronoi_result_median_distance(
*_args,
get_sorted: Callable,
**_kwargs,
) -> torch.Tensor:
return get_sorted().median(dim=-1).values
@staticmethod
def _voronoi_result_softmin(
d: torch.Tensor,
*_args,
temperature=50.0,
use_sorted=None,
d_orig: torch.Tensor,
get_sorted: Callable,
**_kwargs,
) -> torch.Tensor:
d_norm = d_orig.norm(dim=-1)
soft_weights = F.softmax(-d_norm * float(temperature), dim=-1)
eff_d = get_sorted() if use_sorted is not None else d
return (eff_d * soft_weights).sum(dim=-1)
def _voronoi_result_gradient_magnitude(
self,
*args,
name1: str = "f4",
name2: str = "f4",
pad_mode: str = "replicate",
**kwargs,
) -> torch.Tensor:
r1 = self.voronoi_call_mode(name1, result=True, args=args, kwargs=kwargs)
r1_padded = F.pad(r1, (1, 1, 1, 1), mode=pad_mode)
if name2 != name1:
r2 = self.voronoi_call_mode(name2, result=True, args=args, kwargs=kwargs)
r2_padded = F.pad(r2, (1, 1, 1, 1), mode=pad_mode)
else:
r2 = r1
r2_padded = r1_padded
dx = r1_padded[..., 1:-1, 2:] - r2_padded[..., 1:-1, :-2]
dy = r1_padded[..., 2:, 1:-1] - r2_padded[..., :-2, 1:-1]
return (dx**2 + dy**2).sqrt_()
def _voronoi_result_fractal_norm(
self,
d: torch.Tensor,
*args,
name: str = "diff",
mode: str = "sin",
scale: str | float = 0.1,
multiplier: str | float = 10.0,
**kwargs,
) -> torch.Tensor:
if mode == "sin":
fun = torch.sin
elif mode == "cos":
fun = torch.cos
else:
raise ValueError(
"Bad mode parameter for fractal_norm result mode, must be one of: sin, cos",
)
d_adjusted = float(scale) * fun(d * float(multiplier))
my_d_sorted = None
def my_get_sorted():
nonlocal my_d_sorted
if my_d_sorted is not None:
return my_d_sorted
my_d_sorted = d_adjusted.sort(dim=-1).values
return my_d_sorted
return self.voronoi_call_mode(
name,
result=True,
args=(d_adjusted, *args),
kwargs=kwargs | {"get_sorted": my_get_sorted},
)
def _voronoi_result_fuzz(
self,
*args,
name: str = "f1",
fuzz: float | str = 0.25,
**kwargs,
) -> torch.Tensor:
fuzz = float(fuzz)
result = self.voronoi_call_mode(name, result=True, args=args, kwargs=kwargs)
rmin, rmax = result.aminmax()
fuzz = max(abs(rmin.item()), abs(rmax.item())) * fuzz
result += (
torch.rand(result.shape, device=self.gen_device, dtype=result.dtype)
.mul_(fuzz * 2)
.sub_(fuzz)
.to(device=result.device)
)
return utils.normalize_to_scale(result, rmin.item(), rmax.item(), dim=(-2, -1))
def voronoi_distance(self, d: torch.Tensor, octave: int) -> torch.Tensor:
modes = self.get_distance_mode(octave).split("+")
result_scale_base = 1.0 / len(modes)
result = None
for mode in modes:
if ":" in mode:
mode_name, *mode_rest = mode.split(":")
mode_kwargs = dict(
tuple(val.strip() for val in di.split("=", 1)) for di in mode_rest
)
result_scale = result_scale_base * float(mode_kwargs.pop("dscale", 1.0))
else:
mode_name = mode
mode_kwargs = {}
result_scale = result_scale_base
curr_result = self.voronoi_call_mode(
mode_name,
result=False,
args=(d,),
kwargs=mode_kwargs,
).mul_(result_scale)
result = curr_result if result is None else result.add_(curr_result)
return result
def voronoi_result(
self,
d: torch.Tensor,
d_orig: torch.Tensor,
*,
octave: int,
) -> torch.Tensor:
modes = self.get_result_mode(octave).split("+")
result_scale_base = 1.0 / len(modes)
result = None
d_sorted = None
def get_sorted():
nonlocal d_sorted
if d_sorted is not None:
return d_sorted
d_sorted = d.sort(dim=-1).values
return d_sorted
base_kwargs = {
"d_orig": d_orig,
"get_sorted": get_sorted,
}
for mode in modes:
if ":" in mode:
mode_name, *mode_rest = mode.split(":")
mode_kwargs = dict(
tuple(v.strip() for v in di.split("=", 1)) for di in mode_rest
)
result_scale = result_scale_base * float(mode_kwargs.pop("rscale", 1.0))
else:
result_scale = result_scale_base
mode_name = mode
mode_kwargs = {}
curr_result = self.voronoi_call_mode(
mode_name,
result=True,
args=(d,),
kwargs=mode_kwargs | base_kwargs,
).mul_(result_scale)
result = curr_result if result is None else result.add_(curr_result)
return result
def generate_octave(
self,
*,
octave: int,
grid: torch.Tensor,
z_grid: torch.Tensor,
scale: float = 1.0,
) -> torch.Tensor:
# Full 3D grid (H, W, 3)
grid_3d = torch.cat((grid, z_grid), dim=-1)[None, None, ...] # (1, 1, H, W, 3)
grid_3d = grid_3d.expand(self.batch, self.channels, -1, -1, -1)
grid_3d = grid_3d.unsqueeze(-2) # (B, C, H, W, 1, 3)
grid_3d = (grid_3d * scale) % 1.0
# Normalize feature points: already assumed in [0, 1)
fp = self.get_feature_points(octave) # (B, C, N, 3)
fp = fp[:, :, None, None] # (B, C, 1, 1, N, 3)
fp = (fp * scale) % 1.0
# Toroidal wrapped difference
d_orig = d = (grid_3d - fp + 0.5) % 1.0 - 0.5 # Wrap to [-0.5, 0.5)
d = self.voronoi_distance(d.clone(), octave=octave)
return self.voronoi_result(d, d_orig, octave=octave)
def generate(self, *args):
if self.grid_xyz is None or self.feature_points is None or self.z_max == 0:
self.voronoi_reset(*args)
elif self.z_max != 0 and abs(self.z_initial - self.z_curr) > abs(self.z_max):
if self.z_max_mode == "reset":
self.voronoi_reset(*args)
elif self.z_max_mode == "bounce":
self.z_increment = -self.z_increment
self.z_curr += self.z_increment
else:
self.curr_z = self.z_initial
z_range = utils.fallback(self.z_range, max(self.height, self.width))
z_norm = (self.z_curr % z_range) / z_range
self.z_curr += self.z_increment
grid = self.grid_xyz
z_grid = grid.new_full((self.height, self.width, 1), z_norm)
result = grid.new_zeros(self.shape)
amplitude = self.initial_amplitude
scale = self.initial_scale
total_amplitude = 0.0
for octave in range(self.octaves):
result += self.generate_octave(
octave=octave,
grid=grid,
z_grid=z_grid,
scale=scale,
).mul_(amplitude)
total_amplitude += abs(amplitude)
amplitude *= self.gain
scale *= self.lacunarity
result /= total_amplitude if total_amplitude != 0 else 1.0
return result
@@ -0,0 +1,138 @@
# ruff: noqa: ANN002, ANN003
from __future__ import annotations
import torch
from ..utils import fallback
from ..wavelet_functions import Wavelet, wavelet_blend, wavelet_scaling
from .base import FramesToChannelsNoiseGenerator
F = torch.nn.functional
# Idea from https://github.com/ClownsharkBatwing/RES4LYF/ (wave and mode defaults also from that source)
class WaveletFilteredNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "waveletfilter"
MIN_DIMS = 4
MAX_DIMS = 5
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
inv_kwargs = {
k: self.options[k]
for k in ("inv_mode", "inv_biort", "inv_qshift", "inv_wave")
if k in self.options
}
self.wavelet = Wavelet(
wave=self.wave,
level=self.level,
mode=self.mode,
use_1d_dwt=self.use_1d_dwt,
use_dtcwt=self.use_dtcwt,
biort=self.biort,
qshift=self.qshift,
device=self.gen_device,
**inv_kwargs,
)
@classmethod
def ng_params(cls):
return super().ng_params() | {
"mode": "periodization",
"level": 3,
"wave": "haar",
"use_1d_dwt": False,
"use_dtcwt": False,
"qshift": "qshift_a",
"biort": "near_sym_a",
"yl_scale": 1.0,
"yh_scales": 1.0,
"two_step_inverse": False,
"preblend_yl_scale_low": None,
"preblend_yh_scales_low": None,
"preblend_yl_scale_high": None,
"preblend_yh_scales_high": None,
"yl_blend_function": torch.lerp,
"yh_blend_function": torch.lerp,
"yl_blend_high": 0.0,
"yh_blend_high": 1.0,
"noise_sampler": None,
"noise_sampler_high": None,
}
def _fix_shape(self, noise, adjusted_shape):
if noise.shape != adjusted_shape:
noise = noise.reshape(*adjusted_shape)
if self.frames:
noise = noise.reshape(
self.batch,
self.channels * self.frames,
self.height,
self.width,
)
return noise
def generate(self, *args):
adjusted_shape = self.get_adjusted_shape()
noise = (
self.rand_like()
if self.noise_sampler is None
else self.noise_sampler(*args)
)
if self.noise_sampler_high is not None:
noise_high = self._fix_shape(self.noise_sampler_high(*args), adjusted_shape)
else:
noise_high = None
noise = self._fix_shape(noise, adjusted_shape)
orig_noise_shape = noise.shape
need_flat = not self.use_dtcwt and self.use_1d_dwt and noise.ndim > 3
if need_flat:
noise = noise.flatten(start_dim=2)
if noise_high is not None:
noise_high = noise_high.flatten(start_dim=2)
yl, yh = self.wavelet.forward(noise)
if noise_high is not None:
yl_high, yh_high = self.wavelet.forward(noise_high)
if (
self.preblend_yl_scale_high is not None
or self.preblend_yh_scales_high is not None
):
yl_high, yh_high = wavelet_scaling(
yl_high,
yh_high,
fallback(self.preblend_yl_scale_high, 1.0),
fallback(self.preblend_yh_scales_high, 1.0),
)
if (
self.preblend_yl_scale_low is not None
or self.preblend_yh_scales_low is not None
):
yl, yh = wavelet_scaling(
yl,
yh,
fallback(self.preblend_yl_scale_low, 1.0),
fallback(self.preblend_yh_scales_low, 1.0),
)
yl, yh = wavelet_blend(
(yl, yh),
(yl_high, yh_high),
yl_factor=self.yl_blend_high,
yh_factor=self.yh_blend_high,
blend_function=self.yl_blend_function,
yh_blend_function=self.yh_blend_function,
)
del noise_high, yl_high, yh_high
yl, yh = wavelet_scaling(
yl,
yh,
self.yl_scale,
self.yh_scales,
in_place=True,
)
result = self.wavelet.inverse(yl, yh, two_step_inverse=self.two_step_inverse)
if need_flat:
result = result.reshape(orig_noise_shape)
result = self.fix_output_frames(result)
if result.shape == noise.shape:
return result
return result[tuple(slice(0, dl) for dl in noise.shape)]
@@ -0,0 +1,146 @@
# Some noise generation functions shamelessly yoinked from https://github.com/Clybius/ComfyUI-Extra-Samplers
from __future__ import annotations
from typing import TYPE_CHECKING, Any, NamedTuple
import torch
from .. import utils
from .base import FramesToChannelsNoiseGenerator
if TYPE_CHECKING:
from collections.abc import Sequence
class WaveletNoiseOctave(NamedTuple):
octave: int
height: int
width: int
amplitude: float
total_amplitude: float
class WaveletNoiseGenerator(FramesToChannelsNoiseGenerator):
name = "wavelet"
MIN_DIMS = 4
MAX_DIMS = 5
@classmethod
def ng_params(cls):
return super().ng_params() | {
"octave_scale_mode": "adaptive_avg_pool2d",
"octave_rescale_mode": "bilinear",
"post_octave_rescale_mode": "bilinear",
"initial_amplitude": 1.0,
"persistence": 0.5,
"octaves": 4,
"octave_height_factor": 0.5,
"octave_width_factor": 0.5,
"height_factor": 2.0,
"width_factor": 2.0,
"min_height": 4,
"min_width": 4,
"update_blend": 1.0,
"update_blend_function": torch.lerp,
"noise_sampler": None,
}
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
self.set_octave_data()
def set_internal_noise_sampler(self, noise_sampler: object) -> None:
self.noise_sampler = noise_sampler
def set_octave_data(self) -> tuple:
adjusted_shape = self.get_adjusted_shape()
height, width = adjusted_shape[-2:]
amplitude = self.initial_amplitude
total_amplitude = 0.0
curr_height, curr_width = height, width
octave_data = []
is_reverse = self.octaves < 0
octaves = (
range(self.octaves)
if not is_reverse
else reversed(range(abs(self.octaves)))
)
for octave in octaves:
curr_height /= self.height_factor**octave
curr_width /= self.width_factor**octave
if (
amplitude == 0
or curr_height < self.min_height
or curr_width < self.min_width
or curr_height * self.octave_height_factor < 1
or curr_width * self.octave_width_factor < 1
):
if is_reverse and not octave_data:
curr_height, curr_width = height, width
continue
break
total_amplitude += abs(amplitude)
octave_data.append(
WaveletNoiseOctave(
octave=octave,
height=curr_height,
width=curr_width,
amplitude=amplitude,
total_amplitude=total_amplitude,
),
)
amplitude *= self.persistence
if not octave_data or not total_amplitude:
raise ValueError("Unworkable parameters for wavelet noise")
self.octave_data = tuple(octave_data)
def _generate_octave(self, *args: Any, shape: Sequence) -> torch.Tensor:
height, width = shape[-2:]
noise = (
self.noise_sampler(*args)[..., :height, :width].reshape(shape)
if self.noise_sampler
else self.rand_like(shape=(*shape[:-2], height, width))
)
scaled_height = int(max(1, height * self.octave_height_factor))
scaled_width = int(max(1, width * self.octave_width_factor))
scaled_noise = utils.scale_samples(
utils.scale_samples(
noise,
scaled_width,
scaled_height,
mode=self.octave_scale_mode,
),
width=width,
height=height,
mode=self.octave_rescale_mode,
)
return self.update_blend_function(
noise,
noise - scaled_noise,
self.update_blend,
)
def generate(self, *args: Any) -> torch.Tensor:
adjusted_shape = self.get_adjusted_shape()
height, width = adjusted_shape[-2:]
curr_shape = list(adjusted_shape)
result = torch.zeros(
adjusted_shape,
device=self.device,
dtype=self.dtype,
layout=self.layout,
)
for od in self.octave_data:
curr_shape[-2:] = (int(od.height), int(od.width))
octave_output = self._generate_octave(*args, shape=curr_shape)
if octave_output.shape != result.shape:
octave_output = utils.scale_samples(
octave_output,
width,
height,
mode=self.post_octave_rescale_mode,
)
result += octave_output.mul_(od.amplitude)
if self.octave_data[-1].total_amplitude != 0:
result /= self.octave_data[-1].total_amplitude
return self.fix_output_frames(result)
+13 -14
View File
@@ -247,7 +247,7 @@ class SonarBase:
momentum = self.cfg.momentum if momentum is None else momentum
mode = self.cfg.momentum_mode
if (
momentum == 1 # noqa: PLR0916
momentum == 1
or history is None
or (mode == MomentumMode.DENOISED and not is_denoised)
or (mode != MomentumMode.DENOISED and is_denoised)
@@ -412,7 +412,7 @@ class SonarGuidanceMixin:
class SonarWithGuidance(SonarBase, SonarGuidanceMixin):
def __init__(self, *args: list[Any], **kwargs: dict[str, Any]):
def __init__(self, *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)
SonarGuidanceMixin.__init__(self, self.cfg.guidance)
@@ -424,8 +424,8 @@ class SonarSampler(SonarWithGuidance):
sigmas: Tensor,
s_in: Tensor,
extra_args: dict[str, Any],
*args: list[Any],
**kwargs: dict[str, Any],
*args: Any,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
self.model = model
@@ -437,7 +437,7 @@ class SonarSampler(SonarWithGuidance):
self,
x: Tensor,
sigma: Tensor,
*args: list[Any],
*args: Any,
s_in=None,
extra_args=None,
) -> Tensor:
@@ -452,8 +452,8 @@ class SonarSampler(SonarWithGuidance):
class SonarEuler(SonarSampler):
def __init__(
self,
*args: list[Any],
**kwargs: dict[str, Any],
*args: Any,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
@@ -531,8 +531,8 @@ class SonarEulerAncestral(SonarSampler):
self,
eta: float = 1.0,
s_noise: float = 1.0,
*args: list[Any],
**kwargs: dict[str, Any],
*args: Any,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
self.eta = eta
@@ -560,9 +560,8 @@ class SonarEulerAncestral(SonarSampler):
)
if sigma_next > 0:
result_sample = self.guidance_step(step_index, result_sample, denoised)
result_sample = ( # noqa: PLR6104
result_sample
+ self.noise_sampler(sigma, sigma_next) * (self.s_noise * sigma_up)
result_sample = result_sample + self.noise_sampler(sigma, sigma_next) * (
self.s_noise * sigma_up
)
return (
@@ -630,8 +629,8 @@ class SonarDPMPPSDE(SonarSampler):
self,
eta: float = 1.0,
s_noise: float = 1.0,
*args: list[Any],
**kwargs: dict[str, Any],
*args: Any,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
self.eta = eta
+922 -19
View File
File diff suppressed because it is too large Load Diff
+5 -1
View File
@@ -172,6 +172,7 @@ class WCFGPercentages(NamedTuple):
else:
pct_enabled_sigmas = (start_sigma - sigma) / (start_sigma - end_sigma)
steps = len(sigmas) - 1
have_steps = False
if steps > 1:
step = utils.step_from_sigmas(sigma, sigmas)
pct_steps = step / (steps - 1) if step is not None else None
@@ -179,10 +180,11 @@ class WCFGPercentages(NamedTuple):
(sigmas <= start_sigma) & (sigmas >= end_sigma)
]
if len(enabled_steps) > 1:
have_steps = True
step_first = enabled_steps[0].item()
step_last = enabled_steps[-1].item()
pct_enabled_steps = (step - step_first) / (step_last - step_first)
else:
if not have_steps:
step = 0.0
pct_steps = 1.0
step_first = step_last = None
@@ -247,6 +249,8 @@ class WCFGScales(NamedTuple):
target = self.yh_scales
if isinstance(target, float):
return f"{target:.4f}"
if not isinstance(target, (list, tuple)):
return str(target)
result = ", ".join(
self.pretty_yh_scales(target=val)
if isinstance(val, (list, tuple))
+9 -7
View File
@@ -1,13 +1,13 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Callable
from typing import TYPE_CHECKING, Any
import torch
from .utils import fallback
if TYPE_CHECKING:
from collections.abc import Sequence
from collections.abc import Callable, Sequence
try:
import pytorch_wavelets as ptwav
@@ -98,13 +98,15 @@ class Wavelet:
if not two_step_inverse:
return inverse_function((yl, yh))
result = inverse_function((torch.zeros_like(yl), yh))
result += inverse_function((
yl,
tuple(torch.zeros_like(yh_band) for yh_band in yh),
))
result += inverse_function(
(
yl,
tuple(torch.zeros_like(yh_band) for yh_band in yh),
),
)
return result
def to(self, *args: list, copy: bool = False, **kwargs: dict) -> Wavelet:
def to(self, *args: Any, copy: bool = False, **kwargs: Any) -> Wavelet:
o = Wavelet.__new__(Wavelet) if copy else self
o._wavelet_forward = self._wavelet_forward.to(*args, **kwargs) # noqa: SLF001
o._wavelet_inverse = self._wavelet_inverse.to(*args, **kwargs) # noqa: SLF001
+2
View File
@@ -7,6 +7,7 @@ ignore = [
"ANN202",
"ANN204",
"ANN206",
"ANN401",
"C901",
"CPY001",
"DOC201",
@@ -18,6 +19,7 @@ ignore = [
"D105",
"D106",
"D107",
"D401",
"D211",
"D213",
"E402",