diff --git a/README.md b/README.md
index 3c5228a..2768471 100644
--- a/README.md
+++ b/README.md
@@ -64,6 +64,47 @@ You can optionally plug this into the Sonar sampler nodes. See the [Guidance](#g
Very abbreviated section. The init type can make a big difference. If you use `RANDOM` you can get away with setting `direction` to high values (like up to `2.25` or so) and absurdly low values (like `-30.0`). It's also possible to set `momentum` and `momentum_hist` to negative values, although whether it's a good idea...
+
+Click to expand advanced parameters info
+
+There are some extra advanced parameters that may be passed by YAML/JSON using `SamplerConfigOVerride`'s `yaml_parameters`. Defaults:
+
+```yaml
+sonar_params:
+ # One of: classic, new, denoised
+ # classic: Should be the same as the way it works in the A1111 extension.
+ # new: Possibly improved version that doesn't blend in the history again.
+ # denoised: Instead of using the noise prediction, we do momentum on denoised instead.
+ momentum_mode: new
+
+ # The following two parameters may be used to control when
+ # momentum sampling is active. Steps are 0-based with 0 being the first step.
+ momentum_start_step: 0
+ momentum_end_step: 9999
+
+ # Controls whether history always gets updated, whether or not within the
+ # start/end step range or only in that range. Can be used to affect the initial
+ # history value.
+ always_update_history: true
+
+ # Only applies when the init type is RAND.
+ rand_init_noise_multiplier: 1.0
+
+ # If you have ComfyUI-bleh installed, you can use any blend mode it provides.
+ # Otherwise you can have your blend mode in any color you want as long as it's lerp.
+ blend_mode: lerp
+
+ # Defaultss to blend_mode if unset.
+ momentum_blend_mode: null
+
+ # Defaults to blend_mode if unset. Only applies to linear guidance mode.
+ guidance_blend_mode: null
+```
+
+Additionally, it's possible to override the normal Sonar parameters here as well. If they exist in the `sonar_params` block, they will overwrite the values in the node.
+
+
+
## Guidance
You can try the `SamplerSonarNaive` sampler which has an optional latent input. The guidance _probably_ isn't working correctly and the implementation definitely isn't exactly the same as the original A1111 version but it still might be fun to play with. The `linear` guidance type is a lot more sensitive to the `guidance_factor` than the `euler` type. For `euler`, reasonable values are around `0.01` to `0.1`, for `linear` reasonable values are more like `0.001` to `0.02`. It is also possible to set guidance factor to a negative value, I've found this results in high contrast and very vivid colors.
diff --git a/changelog.md b/changelog.md
index 220123a..0621ed5 100644
--- a/changelog.md
+++ b/changelog.md
@@ -2,6 +2,22 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
+## 20241219
+
+*Note*: May change seeds.
+
+This set of changes includes some pretty major internal refactoring. Definitely possible that I broke something, so please create an issue if you run into problems.
+
+* Noise generation should now respect whether generating on CPU vs GPU is selected. Previously it likely was defaulting to generating on GPU. This may change seeds.
+* Refactored momentum samplers, this may change seeds especially if you were using weird parameters like negative direction.
+* Added some new parameters for momentum samplers.
+* Removed the `s_noise` and churn parameters from the normal Sonar Euler sampler. May break workflows. (Churn was the predecessor to ancestral samplers and is basically obsolete.)
+* Added `wavelet` and `distro` noise types.
+* Added `SonarCustomNoiseAdv` node that allows passing parameters via YAML.
+* Added `SonarResizedNoise` node that allows you to generate noise at a fixed size and then crop/resize it to match the generation.
+* Added `SonarAdvancedDistroNoise` node that allows generating noise with basically all the distributions PyTorch supports.
+* Added `SonarWaveletFilteredNoise` node that lets you filter another noise generator using wavelets.
+
## 20241129
*Note*: Contains some potentially workflow-breaking changes.
diff --git a/docs/advanced_noise_nodes.md b/docs/advanced_noise_nodes.md
index 079f4a3..b7b8c11 100644
--- a/docs/advanced_noise_nodes.md
+++ b/docs/advanced_noise_nodes.md
@@ -52,6 +52,21 @@ Parameters:
***
+### `SonarCustomNoiseAdv`
+
+Same as the `SonarCustomNoise` except it also includes a text widget for passing parameters by YAML or JSON (JSON is valid YAML).
+
+Just for example, instead of using the absurdly large `SonarAdvancedDistroNoise` node, you could do something like:
+
+```yaml
+distro: wishart
+quantile_norm: 0.5
+wishart_cov_size: 4
+wishart_df: 3.5
+```
+
+***
+
### `NoisyLatentLike`
This node takes a reference latent and generates noise of the same shape. The one required input is `latent`.
@@ -108,6 +123,59 @@ More extensive documentation TBD (hopefully). For now, a few recipes:
***
+## `SonarWaveletFilteredNoise`
+
+You will need [pytorch_wavelets](https://github.com/fbcotter/pytorch_wavelets) installed in your Python environment to use this one.
+
+Allows filtering another noise source using wavelets. Parameters are specified using YAML (or JSON) in the text widget. The defaults are:
+
+```yaml
+use_dtcwt: false
+mode: periodization
+level: 3
+wave: haar
+
+# Only used in DTCWT mode.
+qshift: qshift_a
+# Only used in DTCWT mode.
+biort: near_sym_a
+
+# Additional parameters for the inverse wavelet operation
+# are null by default and will use whatever the
+# forward parameter is set to:
+# inv_mode, inv_wave, inv_biort, inv_qshift
+# Note: Using different parameters for the inverse wavelet
+# operation may not work well (or just fail entirely).
+
+# Scale for the lowpass filter.
+yl_scale: 1.0
+
+# Scales for the highpass filter. Can be a single value (null is basically 1.0).
+yh_scales: null
+```
+
+***
+
+### `SonarAdvancedDistroNoise`
+
+See: https://pytorch.org/docs/stable/distributions.html
+
+For the most part, we just pass parameters directly to PyTorch's distribution classes. Some of them have specific requirements so it is possible to set invalid parameters.
+
+It may be more convenient to specify parameters using the `SonarCustomNoiseAdv` node than this gigantic monstrosity of a node. **Note**: In that case, pass the distribution name using `distro`, i.e. `distro: laplacian`.
+
+Common parameters:
+
+* `quantile_norm`: When enabled, will normalize generated noise to this quantile (i.e. 0.75 means outliers >75% will be clipped). Set to 1.0 or 0.0 to disable quantile normalization. A value like 0.75 or 0.85 should be reasonable, it really depends on the distribution and how many of the values are extreme. Some actually work better with quantile normalization disabled.
+* `quantile_norm_mode`: Controls what dimensions quantile normalization uses. By default, the noise is flattened first. You can try the nonflat versions but they may have a very strong row/column influence. Only applies when quantile_norm is active.
+* `result_index`: When noise generation returns a batch of items, it will select the specified index. Negative indexes count from the end. Values outside the valid range will be automatically adjusted. You may enter a space-separated list of values for the case where there might be multiple added batch dimensions. Excess batch dimensions are removed from the end, indexe from result_index are used in order so you may want to enter the indexes in reverse order. Example: If your noise has shape `(1, 4, 3, 3)` and two 2-sized batch dims are added resulting in `(1, 4, 3, 3, 2, 2)` and you wanted index 0 from the first additional batch dimension and 1 from the second you would use result_index: `1 0`
+
+Individual distributions have parameters beginning with their name, i.e. `laplacian_loc`. Parameters that are string inputs usually allow entering multiple space-separated items. This will usually result in the output noise being a batch, which can be selected with the `result_index` parameter.
+
+Suggestions for fun distributions to try: Wishart and VonMises can produce some interesting results.
+
+***
+
### `SonarModulatedNoise`
Experimental noise modulation based on code stolen from
diff --git a/py/external.py b/py/external.py
index 4d2474c..1faed51 100644
--- a/py/external.py
+++ b/py/external.py
@@ -1,22 +1,125 @@
+from __future__ import annotations
+
import contextlib
import importlib
+import sys
+from functools import partial
+from typing import TYPE_CHECKING, Callable, NamedTuple
-MODULES = {}
+if TYPE_CHECKING:
+ from types import ModuleType
-with contextlib.suppress(ImportError, NotImplementedError):
- bleh = importlib.import_module("custom_nodes.ComfyUI-bleh")
- bleh_version = getattr(bleh, "BLEH_VERSION", -1)
- if bleh_version < 1:
- raise NotImplementedError
- MODULES["bleh"] = bleh
-with contextlib.suppress(ImportError, NotImplementedError):
- import custom_nodes.ComfyUI_restart_sampling as rs
+class Integrations:
+ class Integration(NamedTuple):
+ key: str
+ module_name: str
+ handler: Callable | None = None
+
+ def __init__(self):
+ self.initialized = False
+ self.modules = {}
+ self.init_handlers = []
+ self.handlers = []
+
+ def __getitem__(self, key):
+ return self.modules[key]
+
+ def __contains__(self, key):
+ return key in self.modules
+
+ def __getattr__(self, key):
+ return self.modules.get(key)
+
+ @staticmethod
+ def get_custom_node(name: str) -> ModuleType | None:
+ module_key = f"custom_nodes.{name}"
+ with contextlib.suppress(StopIteration):
+ spec = importlib.util.find_spec(module_key)
+ if spec is None:
+ return None
+ return next(
+ v
+ for v in sys.modules.copy().values()
+ if hasattr(v, "__spec__")
+ and v.__spec__ is not None
+ and v.__spec__.origin == spec.origin
+ )
+ return None
+
+ def register_init_handler(self, handler):
+ self.init_handlers.append(handler)
+
+ def register_integration(self, key: str, module_name: str, handler=None) -> None:
+ if self.initialized:
+ raise ValueError(
+ "Internal error: Cannot register integration after initialization",
+ )
+ if any(item[0] == key or item[1] == module_name for item in self.handlers):
+ errstr = (
+ f"Module {module_name} ({key}) already in integration handlers list!"
+ )
+ raise ValueError(errstr)
+ self.handlers.append(self.Integration(key, module_name, handler))
+
+ def initialize(self) -> None:
+ if self.initialized:
+ return
+ self.initialized = True
+ for ih in self.handlers:
+ module = self.get_custom_node(ih.module_name)
+ if module is None:
+ continue
+ if ih.handler is not None:
+ module = ih.handler(module)
+ if module is not None:
+ self.modules[ih.key] = module
+
+ for init_handler in self.init_handlers:
+ init_handler(self)
+
+
+class SonarIntegrations(Integrations):
+ def __init__(self, *args: list, **kwargs: dict):
+ super().__init__(*args, **kwargs)
+ self.register_integration("bleh", "ComfyUI-bleh", self.bleh_integration)
+ self.register_integration(
+ "restart",
+ "ComfyUI_restart_sampling",
+ self.restart_integration,
+ )
+
+ @classmethod
+ def bleh_integration(cls, module: ModuleType) -> ModuleType | None:
+ bleh_version = getattr(module, "BLEH_VERSION", -1)
+ if bleh_version < 1:
+ return None
+ return module
+
+ @classmethod
+ def restart_integration(cls, module: ModuleType) -> ModuleType | None:
+ if hasattr(module, "restart_sampling") and hasattr(
+ module.restart_sampling,
+ "DEFAULT_SEGMENTS",
+ ):
+ return module
+ return None
+
+
+MODULES = SonarIntegrations()
+
+
+class IntegratedNode(type):
+ @staticmethod
+ def wrap_INPUT_TYPES(orig_method: Callable, *args: list, **kwargs: dict) -> dict:
+ MODULES.initialize()
+ return orig_method(*args, **kwargs)
+
+ def __new__(cls: type, name: str, bases: tuple, attrs: dict) -> object:
+ obj = type.__new__(cls, name, bases, attrs)
+ if hasattr(obj, "INPUT_TYPES"):
+ obj.INPUT_TYPES = partial(cls.wrap_INPUT_TYPES, obj.INPUT_TYPES)
+ return obj
- if not hasattr(rs.restart_sampling, "DEFAULT_SEGMENTS"):
- # Dumb test but this should only exist in restart sampling versions that
- # support plugging in custom noise.
- raise NotImplementedError
- MODULES["restart"] = rs
__all__ = ("MODULES",)
diff --git a/py/freeu_extreme.py b/py/freeu_extreme.py
index d9b401d..d94ea40 100644
--- a/py/freeu_extreme.py
+++ b/py/freeu_extreme.py
@@ -2,7 +2,8 @@ from __future__ import annotations
import torch
-from .external import MODULES as EXTERNAL_MODULES
+from . import utils
+from .external import IntegratedNode
from .powernoise import PowerFilter
@@ -28,14 +29,7 @@ def ffilter(x, pfilter, normalization_factor=1.0, cfg_idx=None, filter_cache=Non
return x_filt.to(x.dtype, non_blocking=True)
-BLEND_OPS = (
- {"lerp": torch.lerp}
- if "bleh" not in EXTERNAL_MODULES
- else EXTERNAL_MODULES["bleh"].py.latent_utils.BLENDING_MODES
-)
-
-
-class FreeUExtremeConfigNode:
+class FreeUExtremeConfigNode(metaclass=IntegratedNode):
DESCRIPTION = "Allows setting configuration for FreeU Extreme."
RETURN_TYPES = ("FRUX_CONFIG",)
FUNCTION = "go"
@@ -150,7 +144,7 @@ class FreeUExtremeConfigNode:
},
),
"blend_mode": (
- tuple(BLEND_OPS.keys()),
+ tuple(utils.BLENDING_MODES.keys()),
{
"tooltip": "Mode used when blending. Generally only has an effect when blend is set to values other than 0 or 1",
},
@@ -302,7 +296,7 @@ class FreeUExtremeConfig:
x[:, slice_offs : slice_offs + slice_size] = (
xslice
if self.blend == 1.0
- else BLEND_OPS[self.blend_mode](
+ else utils.BLENDING_MODES[self.blend_mode](
x[:, slice_offs : slice_offs + slice_size],
xslice,
self.blend,
@@ -331,12 +325,12 @@ class FreeUExtremeConfig:
def clone(self):
return self.__class__(**{k: getattr(self, k) for k in self._keys})
- def __repr__(self): # noqa: D105
+ def __repr__(self):
meh = {k: getattr(self, k) for k in self._keys}
return f""
-class FreeUExtremeNode:
+class FreeUExtremeNode(metaclass=IntegratedNode):
DESCRIPTION = "Main FreeU Extreme node. Allows patching a model with the FreeU (V2) effect with more control."
RETURN_TYPES = ("MODEL",)
FUNCTION = "go"
diff --git a/py/nodes.py b/py/nodes.py
index ef08961..e6637ce 100644
--- a/py/nodes.py
+++ b/py/nodes.py
@@ -11,9 +11,9 @@ import torch
import yaml
from comfy import model_management, samplers
-from . import external, noise
+from . import external, noise, utils
+from .external import IntegratedNode
from .noise import NoiseType
-from .noise_utils import scale_noise
from .sonar import (
GuidanceConfig,
GuidanceType,
@@ -52,7 +52,7 @@ if not HAVE_COMFY_UNION_TYPE:
result.whitelist = whitelist
return result
- def __ne__(self, other): # noqa: D105
+ def __ne__(self, other):
return False if self.whitelist is None else other not in self.whitelist
WILDCARD_NOISE = Wildcard("*", whitelist=NOISE_INPUT_TYPES)
@@ -65,24 +65,7 @@ NOISE_INPUT_TYPES_HINT = (
)
-if "bleh" in external.MODULES:
- bleh_latent_utils = external.MODULES["bleh"].py.latent_utils
- BLEND_MODES = bleh_latent_utils.BLENDING_MODES
- UPSCALE_METHODS = bleh_latent_utils.UPSCALE_METHODS
- del bleh_latent_utils
-else:
- BLEND_MODES = {"lerp": torch.lerp}
- UPSCALE_METHODS = (
- "bilinear",
- "nearest-exact",
- "nearest",
- "area",
- "bicubic",
- "bislerp",
- )
-
-
-class NoisyLatentLikeNode:
+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."
RETURN_TYPES = ("LATENT",)
OUTPUT_TOOLTIPS = ("The noisy latent image.",)
@@ -249,7 +232,7 @@ class NoisyLatentLikeNode:
)
finally:
torch.random.set_rng_state(randst)
- result = scale_noise(result, multiplier, normalized=True)
+ result = utils.scale_noise(result, multiplier, normalized=True)
if add_to_latent:
result += latent_samples.repeat(
*(repeat_batch if i == 0 else 1 for i in range(latent_samples.ndim)),
@@ -258,7 +241,7 @@ class NoisyLatentLikeNode:
return ({"samples": result},)
-class SonarCustomNoiseNodeBase(abc.ABC):
+class SonarCustomNoiseNodeBase(metaclass=IntegratedNode):
DESCRIPTION = "A custom noise item."
RETURN_TYPES = ("SONAR_CUSTOM_NOISE",)
OUTPUT_TOOLTIPS = ("A custom noise chain.",)
@@ -771,7 +754,7 @@ class SonarGuidedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMixi
):
return super().go(
factor,
- ref_latent=scale_noise(
+ ref_latent=utils.scale_noise(
SonarGuidanceMixin.prepare_ref_latent(latent["samples"].clone()),
normalized=normalize_ref,
),
@@ -899,7 +882,7 @@ class SonarBlendedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix
},
),
"blend_mode": (
- tuple(BLEND_MODES.keys()),
+ tuple(utils.BLENDING_MODES.keys()),
{
"default": "lerp",
"tooltip": "Mode used for blending the two noise types. More modes will be available if ComfyUI-bleh is installed.",
@@ -944,7 +927,7 @@ class SonarBlendedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix
custom_noise_2=None,
blend_mode="lerp",
):
- blend_function = BLEND_MODES.get(blend_mode)
+ blend_function = utils.BLENDING_MODES.get(blend_mode)
if blend_function is None:
raise ValueError("Unknown blend mode")
return super().go(
@@ -1032,14 +1015,14 @@ class SonarResizedNoiseNode(SonarCustomNoiseNodeBase, SonarNormalizeNoiseNodeMix
},
),
"upscale_mode": (
- UPSCALE_METHODS,
+ utils.UPSCALE_METHODS,
{
"tooltip": "Allows setting the scaling mode when width/height is smaller than the requested size.",
"default": "nearest-exact",
},
),
"downscale_mode": (
- UPSCALE_METHODS,
+ utils.UPSCALE_METHODS,
{
"tooltip": "Allows setting the scaling mode when width/height is larger than the requested size and downscale_strategy is set to 'scale'.",
"default": "nearest-exact",
@@ -1133,7 +1116,7 @@ class SonarAdvancedPyramidNoiseNode(SonarCustomNoiseNodeBase):
},
),
"upscale_mode": (
- ("default", *UPSCALE_METHODS),
+ ("default", *utils.UPSCALE_METHODS),
{
"tooltip": "Allows setting the scaling mode for Pyramid noise. Leave on default to use the variant default.",
"default": "default",
@@ -1552,7 +1535,7 @@ class SonarWaveletFilteredNoiseNode(
)
-class SonarToComfyNOISENode:
+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)."
RETURN_TYPES = ("NOISE",)
CATEGORY = "sampling/custom_sampling/noise"
@@ -1688,7 +1671,7 @@ class GuidanceConfigNode:
)
-class SamplerNodeSonarBase:
+class SamplerNodeSonarBase(metaclass=IntegratedNode):
DESCRIPTION = "Sonar - momentum based sampler node."
@classmethod
@@ -1953,9 +1936,8 @@ class SamplerNodeSonarDPMPPSDE(SamplerNodeSonarEuler):
)
-class SamplerNodeConfigOverride:
+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."
- # KWARG_OVERRIDES = ("s_noise", "eta", "s_churn", "r", "solver_type")
@classmethod
def INPUT_TYPES(cls):
@@ -2167,6 +2149,8 @@ class SamplerNodeConfigOverride:
)
+NODE_DISPLAY_NAME_MAPPINGS = {}
+
NODE_CLASS_MAPPINGS = {
"SamplerSonarEuler": SamplerNodeSonarEuler,
"SamplerSonarEulerA": SamplerNodeSonarEulerAncestral,
@@ -2193,229 +2177,240 @@ NODE_CLASS_MAPPINGS = {
"SONAR_CUSTOM_NOISE to NOISE": SonarToComfyNOISENode,
}
-NODE_DISPLAY_NAME_MAPPINGS = {}
+
+bleh = None
-if "bleh" in external.MODULES:
- import ast
+class SonarBlendFilterNoiseNode(
+ SonarCustomNoiseNodeBase,
+ SonarNormalizeNoiseNodeMixin,
+):
+ DESCRIPTION = "Custom noise type that allows blending and filtering the output of another noise generator using ComfyUI-bleh."
- bleh = external.MODULES["bleh"]
- bleh_latentutils = bleh.py.latent_utils
- bleh_ops = bleh.py.nodes.ops
+ @classmethod
+ def INPUT_TYPES(cls):
+ result = super().INPUT_TYPES(include_rescale=False, include_chain=False)
+ result["required"] |= {
+ "sonar_custom_noise": (
+ WILDCARD_NOISE,
+ {
+ "tooltip": f"Custom noise input.\n{NOISE_INPUT_TYPES_HINT}",
+ },
+ ),
+ "blend_mode": (("simple_add", *utils.BLENDING_MODES.keys()),),
+ "ffilter": (tuple(bleh.py.latent_utils.FILTER_PRESETS.keys()),),
+ "ffilter_custom": ("STRING", {"default": ""}),
+ "ffilter_scale": (
+ "FLOAT",
+ {"default": 1.0, "min": -100.0, "max": 100.0},
+ ),
+ "ffilter_strength": (
+ "FLOAT",
+ {"default": 0.0, "min": -100.0, "max": 100.0},
+ ),
+ "ffilter_threshold": (
+ "INT",
+ {"default": 1, "min": 1, "max": 32},
+ ),
+ "enhance_mode": (("none", *bleh.py.latent_utils.ENHANCE_METHODS),),
+ "enhance_strength": (
+ "FLOAT",
+ {"default": 0.0, "min": -100.0, "max": 100.0},
+ ),
+ "affect": (("result", "noise", "both"),),
+ "normalize_result": (("default", "forced", "disabled"),),
+ "normalize_noise": (("default", "forced", "disabled"),),
+ }
+ return result
- class SonarBlendFilterNoiseNode(
- SonarCustomNoiseNodeBase,
- SonarNormalizeNoiseNodeMixin,
+ @classmethod
+ def get_item_class(cls):
+ return noise.BlendFilterNoise
+
+ def go(
+ self,
+ *,
+ factor,
+ sonar_custom_noise,
+ blend_mode,
+ ffilter,
+ ffilter_custom,
+ ffilter_scale,
+ ffilter_strength,
+ ffilter_threshold,
+ enhance_mode,
+ enhance_strength,
+ affect,
+ normalize_result,
+ normalize_noise,
):
- DESCRIPTION = "Custom noise type that allows blending and filtering the output of another noise generator using ComfyUI-bleh."
+ import ast # noqa: PLC0415
- @classmethod
- def INPUT_TYPES(cls):
- result = super().INPUT_TYPES(include_rescale=False, include_chain=False)
- result["required"] |= {
- "sonar_custom_noise": (
- WILDCARD_NOISE,
- {
- "tooltip": f"Custom noise input.\n{NOISE_INPUT_TYPES_HINT}",
- },
- ),
- "blend_mode": (
- ("simple_add", *bleh_latentutils.BLENDING_MODES.keys()),
- ),
- "ffilter": (tuple(bleh_latentutils.FILTER_PRESETS.keys()),),
- "ffilter_custom": ("STRING", {"default": ""}),
- "ffilter_scale": (
- "FLOAT",
- {"default": 1.0, "min": -100.0, "max": 100.0},
- ),
- "ffilter_strength": (
- "FLOAT",
- {"default": 0.0, "min": -100.0, "max": 100.0},
- ),
- "ffilter_threshold": (
- "INT",
- {"default": 1, "min": 1, "max": 32},
- ),
- "enhance_mode": (("none", *bleh_latentutils.ENHANCE_METHODS),),
- "enhance_strength": (
- "FLOAT",
- {"default": 0.0, "min": -100.0, "max": 100.0},
- ),
- "affect": (("result", "noise", "both"),),
- "normalize_result": (("default", "forced", "disabled"),),
- "normalize_noise": (("default", "forced", "disabled"),),
- }
- return result
-
- @classmethod
- def get_item_class(cls):
- return noise.BlendFilterNoise
-
- def go(
- self,
- *,
+ ffilter_custom = ffilter_custom.strip()
+ normalize_result = (
+ None if normalize_result == "default" else normalize_result == "forced"
+ )
+ normalize_noise = (
+ None if normalize_noise == "default" else normalize_noise == "forced"
+ )
+ if ffilter_custom:
+ ffilter = ast.literal_eval(f"[{ffilter_custom}]")
+ else:
+ ffilter = bleh.py.latent_utils.FILTER_PRESETS[ffilter]
+ return super().go(
factor,
- sonar_custom_noise,
- blend_mode,
- ffilter,
- ffilter_custom,
- ffilter_scale,
- ffilter_strength,
- ffilter_threshold,
- enhance_mode,
- enhance_strength,
- affect,
- normalize_result,
- normalize_noise,
- ):
- ffilter_custom = ffilter_custom.strip()
- normalize_result = (
- None if normalize_result == "default" else normalize_result == "forced"
- )
- normalize_noise = (
- None if normalize_noise == "default" else normalize_noise == "forced"
- )
- if ffilter_custom:
- ffilter = ast.literal_eval(f"[{ffilter_custom}]")
- else:
- ffilter = bleh_latentutils.FILTER_PRESETS[ffilter]
- return super().go(
- factor,
- noise=sonar_custom_noise.clone(),
- blend_mode=blend_mode,
- ffilter=ffilter,
- ffilter_scale=ffilter_scale,
- ffilter_strength=ffilter_strength,
- ffilter_threshold=ffilter_threshold,
- enhance_mode=enhance_mode,
- enhance_strength=enhance_strength,
- affect=affect,
- normalize_noise=self.get_normalize(normalize_noise),
- normalize_result=self.get_normalize(normalize_result),
- )
-
- class SonarBlehOpsNoiseNode(
- SonarCustomNoiseNodeBase,
- SonarNormalizeNoiseNodeMixin,
- ):
- DESCRIPTION = (
- "Custom noise type that allows manipulating noise with ComfyUI-bleh ops."
+ noise=sonar_custom_noise.clone(),
+ blend_mode=blend_mode,
+ ffilter=ffilter,
+ ffilter_scale=ffilter_scale,
+ ffilter_strength=ffilter_strength,
+ ffilter_threshold=ffilter_threshold,
+ enhance_mode=enhance_mode,
+ enhance_strength=enhance_strength,
+ affect=affect,
+ normalize_noise=self.get_normalize(normalize_noise),
+ normalize_result=self.get_normalize(normalize_result),
)
- @classmethod
- def INPUT_TYPES(cls):
- result = super().INPUT_TYPES(include_rescale=False, include_chain=False)
- result["required"] |= {
- "sonar_custom_noise": (
- WILDCARD_NOISE,
- {
- "tooltip": f"Custom noise input.\n{NOISE_INPUT_TYPES_HINT}",
- },
+
+class SonarBlehOpsNoiseNode(
+ SonarCustomNoiseNodeBase,
+ SonarNormalizeNoiseNodeMixin,
+):
+ DESCRIPTION = (
+ "Custom noise type that allows manipulating noise with ComfyUI-bleh ops."
+ )
+
+ @classmethod
+ def INPUT_TYPES(cls):
+ result = super().INPUT_TYPES(include_rescale=False, include_chain=False)
+ result["required"] |= {
+ "sonar_custom_noise": (
+ WILDCARD_NOISE,
+ {
+ "tooltip": f"Custom noise input.\n{NOISE_INPUT_TYPES_HINT}",
+ },
+ ),
+ "normalize": (
+ ("default", "forced", "disabled"),
+ {
+ "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.",
+ },
+ ),
+ "rules": (
+ "STRING",
+ {
+ "tooltip": "Enter rules in the bleh block ops format here.",
+ "placeholder": "# YAML ops here",
+ "dynamicPrompts": False,
+ "multiline": True,
+ },
+ ),
+ }
+ return result
+
+ @classmethod
+ def get_item_class(cls):
+ return noise.BlehOpsNoise
+
+ def go(
+ self,
+ *,
+ factor,
+ sonar_custom_noise,
+ rules,
+ normalize,
+ ):
+ return super().go(
+ factor,
+ noise=sonar_custom_noise.clone(),
+ rules=bleh.py.nodes.ops.RuleGroup.from_yaml(rules),
+ normalize=normalize,
+ )
+
+
+restart = None
+
+
+class KRestartSamplerCustomNoise(metaclass=IntegratedNode):
+ DESCRIPTION = "Restart sampler variant that allows specifying a custom noise type for noise added by restarts."
+
+ @classmethod
+ def INPUT_TYPES(cls):
+ get_normal_schedulers = getattr(
+ restart.nodes,
+ "get_supported_normal_schedulers",
+ restart.nodes.get_supported_restart_schedulers,
+ )
+ return {
+ "required": {
+ "model": ("MODEL",),
+ "add_noise": (["enable", "disable"],),
+ "noise_seed": (
+ "INT",
+ {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF},
),
- "normalize": (
- ("default", "forced", "disabled"),
- {
- "tooltip": "Controls whether the generated noise is normalized to 1.0 strength.",
- },
- ),
- "rules": (
+ "steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
+ "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
+ "sampler": ("SAMPLER",),
+ "scheduler": (get_normal_schedulers(),),
+ "positive": ("CONDITIONING",),
+ "negative": ("CONDITIONING",),
+ "latent_image": ("LATENT",),
+ "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}),
+ "end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}),
+ "return_with_leftover_noise": (["disable", "enable"],),
+ "segments": (
"STRING",
{
- "tooltip": "Enter rules in the bleh block ops format here.",
- "placeholder": "# YAML ops here",
- "dynamicPrompts": False,
- "multiline": True,
+ "default": restart.restart_sampling.DEFAULT_SEGMENTS,
+ "multiline": False,
},
),
- }
- return result
+ "restart_scheduler": (
+ restart.nodes.get_supported_restart_schedulers(),
+ ),
+ "chunked_mode": ("BOOLEAN", {"default": True}),
+ },
+ "optional": {
+ "custom_noise_opt": (
+ WILDCARD_NOISE,
+ {
+ "tooltip": f"Optional custom noise input.\n{NOISE_INPUT_TYPES_HINT}",
+ },
+ ),
+ },
+ }
- @classmethod
- def get_item_class(cls):
- return noise.BlehOpsNoise
+ RETURN_TYPES = ("LATENT", "LATENT")
+ RETURN_NAMES = ("output", "denoised_output")
+ FUNCTION = "sample"
+ CATEGORY = "sampling"
- def go(
- self,
- *,
- factor,
- sonar_custom_noise,
- rules,
- normalize,
- ):
- return super().go(
- factor,
- noise=sonar_custom_noise.clone(),
- rules=bleh_ops.RuleGroup.from_yaml(rules),
- normalize=normalize,
- )
-
- NODE_CLASS_MAPPINGS |= {
- "SonarBlendFilterNoise": SonarBlendFilterNoiseNode,
- "SonarBlehOpsNoise": SonarBlehOpsNoiseNode,
- }
-
-if "restart" in external.MODULES:
- rs = external.MODULES["restart"]
-
- class KRestartSamplerCustomNoise:
- DESCRIPTION = "Restart sampler variant that allows specifying a custom noise type for noise added by restarts."
-
- @classmethod
- def INPUT_TYPES(cls):
- get_normal_schedulers = getattr(
- rs.nodes,
- "get_supported_normal_schedulers",
- rs.nodes.get_supported_restart_schedulers,
- )
- return {
- "required": {
- "model": ("MODEL",),
- "add_noise": (["enable", "disable"],),
- "noise_seed": (
- "INT",
- {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF},
- ),
- "steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
- "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
- "sampler": ("SAMPLER",),
- "scheduler": (get_normal_schedulers(),),
- "positive": ("CONDITIONING",),
- "negative": ("CONDITIONING",),
- "latent_image": ("LATENT",),
- "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}),
- "end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}),
- "return_with_leftover_noise": (["disable", "enable"],),
- "segments": (
- "STRING",
- {
- "default": rs.restart_sampling.DEFAULT_SEGMENTS,
- "multiline": False,
- },
- ),
- "restart_scheduler": (rs.nodes.get_supported_restart_schedulers(),),
- "chunked_mode": ("BOOLEAN", {"default": True}),
- },
- "optional": {
- "custom_noise_opt": (
- WILDCARD_NOISE,
- {
- "tooltip": f"Optional custom noise input.\n{NOISE_INPUT_TYPES_HINT}",
- },
- ),
- },
- }
-
- RETURN_TYPES = ("LATENT", "LATENT")
- RETURN_NAMES = ("output", "denoised_output")
- FUNCTION = "sample"
- CATEGORY = "sampling"
-
- @classmethod
- def sample(
- cls,
- *,
+ @classmethod
+ def sample(
+ cls,
+ *,
+ model,
+ add_noise,
+ noise_seed,
+ steps,
+ cfg,
+ sampler,
+ scheduler,
+ positive,
+ negative,
+ latent_image,
+ start_at_step,
+ end_at_step,
+ return_with_leftover_noise,
+ segments,
+ restart_scheduler,
+ chunked_mode=False,
+ custom_noise_opt=None,
+ ):
+ return restart.restart_sampling.restart_sampling(
model,
- add_noise,
noise_seed,
steps,
cfg,
@@ -2424,78 +2419,74 @@ if "restart" in external.MODULES:
positive,
negative,
latent_image,
- start_at_step,
- end_at_step,
- return_with_leftover_noise,
segments,
restart_scheduler,
- chunked_mode=False,
- custom_noise_opt=None,
- ):
- return rs.restart_sampling.restart_sampling(
- model,
- noise_seed,
- steps,
- cfg,
- sampler,
- scheduler,
- positive,
- negative,
- latent_image,
- segments,
- restart_scheduler,
- disable_noise=add_noise == "disable",
- step_range=(start_at_step, end_at_step),
- force_full_denoise=return_with_leftover_noise != "enable",
- output_only=False,
- chunked_mode=chunked_mode,
- custom_noise=custom_noise_opt.make_noise_sampler
- if custom_noise_opt
- else None,
- )
+ disable_noise=add_noise == "disable",
+ step_range=(start_at_step, end_at_step),
+ force_full_denoise=return_with_leftover_noise != "enable",
+ output_only=False,
+ chunked_mode=chunked_mode,
+ custom_noise=custom_noise_opt.make_noise_sampler
+ if custom_noise_opt
+ else None,
+ )
- NODE_CLASS_MAPPINGS["KRestartSamplerCustomNoise"] = KRestartSamplerCustomNoise
- if hasattr(rs.restart_sampling, "RestartSampler"):
+class RestartSamplerCustomNoise(metaclass=IntegratedNode):
+ DESCRIPTION = "Wrapper used to make another sampler Restart compatible. Allows specifying a custom type for noise added by restarts."
- class RestartSamplerCustomNoise:
- DESCRIPTION = "Wrapper used to make another sampler Restart compatible. Allows specifying a custom type for noise added by restarts."
-
- @classmethod
- def INPUT_TYPES(cls):
- return {
- "required": {
- "sampler": ("SAMPLER",),
- "chunked_mode": ("BOOLEAN", {"default": True}),
+ @classmethod
+ def INPUT_TYPES(cls):
+ return {
+ "required": {
+ "sampler": ("SAMPLER",),
+ "chunked_mode": ("BOOLEAN", {"default": True}),
+ },
+ "optional": {
+ "custom_noise_opt": (
+ WILDCARD_NOISE,
+ {
+ "tooltip": f"Optional custom noise input.\n{NOISE_INPUT_TYPES_HINT}",
},
- "optional": {
- "custom_noise_opt": (
- WILDCARD_NOISE,
- {
- "tooltip": f"Optional custom noise input.\n{NOISE_INPUT_TYPES_HINT}",
- },
- ),
- },
- }
+ ),
+ },
+ }
- RETURN_TYPES = ("SAMPLER",)
- FUNCTION = "go"
- CATEGORY = "sampling/custom_sampling/samplers"
+ RETURN_TYPES = ("SAMPLER",)
+ FUNCTION = "go"
+ CATEGORY = "sampling/custom_sampling/samplers"
- @classmethod
- def go(cls, sampler, chunked_mode, custom_noise_opt=None):
- restart_options = {
- "restart_chunked": chunked_mode,
- "restart_wrapped_sampler": sampler,
- "restart_custom_noise": None
- if custom_noise_opt is None
- else custom_noise_opt.make_noise_sampler,
- }
- restart_sampler = samplers.KSAMPLER(
- rs.restart_sampling.RestartSampler.sampler_function,
- extra_options=sampler.extra_options | restart_options,
- inpaint_options=sampler.inpaint_options,
- )
- return (restart_sampler,)
+ @classmethod
+ def go(cls, sampler, chunked_mode, custom_noise_opt=None):
+ restart_options = {
+ "restart_chunked": chunked_mode,
+ "restart_wrapped_sampler": sampler,
+ "restart_custom_noise": None
+ if custom_noise_opt is None
+ else custom_noise_opt.make_noise_sampler,
+ }
+ restart_sampler = samplers.KSAMPLER(
+ restart.restart_sampling.RestartSampler.sampler_function,
+ extra_options=sampler.extra_options | restart_options,
+ inpaint_options=sampler.inpaint_options,
+ )
+ return (restart_sampler,)
- NODE_CLASS_MAPPINGS["RestartSamplerCustomNoise"] = RestartSamplerCustomNoise
+
+def init_integrations(integrations):
+ global NODE_CLASS_MAPPINGS, restart, bleh # noqa: PLW0603
+ restart = integrations.restart
+ if restart is not None:
+ NODE_CLASS_MAPPINGS["KRestartSamplerCustomNoise"] = KRestartSamplerCustomNoise
+ if hasattr(restart.restart_sampling, "RestartSampler"):
+ NODE_CLASS_MAPPINGS["RestartSamplerCustomNoise"] = RestartSamplerCustomNoise
+ bleh = integrations.bleh
+ if bleh is None:
+ return
+ NODE_CLASS_MAPPINGS |= {
+ "SonarBlendFilterNoise": SonarBlendFilterNoiseNode,
+ "SonarBlehOpsNoise": SonarBlehOpsNoiseNode,
+ }
+
+
+external.MODULES.register_init_handler(init_integrations)
diff --git a/py/noise.py b/py/noise.py
index fea8139..80bb949 100644
--- a/py/noise.py
+++ b/py/noise.py
@@ -10,10 +10,10 @@ import yaml
from comfy.k_diffusion import sampling
from torch import Tensor
-from . import external
+from . import external, utils
from .noise_generation import *
-from .noise_utils import crop_samples, scale_noise, scale_samples
from .sonar import SonarGuidanceMixin
+from .utils import crop_samples, scale_noise
# ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311
@@ -1070,13 +1070,13 @@ class ResizedNoise(CustomNoiseItemBase):
crop_mode=crop_mode,
upscale_mode=upscale_mode,
downscale_mode=downscale_mode,
- noise=custom_noise.clone(),
+ custom_noise=custom_noise.clone(),
normalize=normalize,
)
def clone_key(self, k):
- if k == "noise":
- return self.noise.clone()
+ if k == "custom_noise":
+ return self.custom_noise.clone()
return super().clone_key(k)
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
@@ -1116,15 +1116,25 @@ class ResizedNoise(CustomNoiseItemBase):
offset_height=offsh,
)
else:
- x = scale_samples(x, nw, nh, mode=self.downscale_mode)
- output = partial(scale_samples, width=xw, height=xh, mode=upscale_mode)
+ x = utils.scale_samples(x, nw, nh, mode=self.downscale_mode)
+ output = partial(
+ utils.scale_samples,
+ width=xw,
+ height=xh,
+ mode=upscale_mode,
+ )
else:
- x = scale_samples(x, nw, nh, mode=self.upscale_mode)
+ x = utils.scale_samples(x, nw, nh, mode=self.upscale_mode)
if x_any_bigger:
- output = partial(scale_samples, width=xw, height=xh, mode=upscale_mode)
+ output = partial(
+ utils.scale_samples,
+ width=xw,
+ height=xh,
+ mode=upscale_mode,
+ )
elif self.downscale_strategy == "scale":
output = partial(
- scale_samples,
+ utils.scale_samples,
width=xw,
height=xh,
mode=downscale_mode,
@@ -1138,7 +1148,7 @@ class ResizedNoise(CustomNoiseItemBase):
offset_width=offsw,
offset_height=offsh,
)
- ns = self.noise.make_noise_sampler(x, *args, normalized=False, **kwargs)
+ ns = self.custom_noise.make_noise_sampler(x, *args, normalized=False, **kwargs)
del x
def noise_sampler(*args, **kwargs):
@@ -1206,160 +1216,158 @@ class WaveletFilteredNoise(CustomNoiseItemBase):
return noise_sampler
-if "bleh" in external.MODULES:
- bleh = external.MODULES["bleh"]
- BLU = bleh.py.latent_utils
- BOPS = bleh.py.nodes.ops
-
- class BlendFilterNoise(CustomNoiseItemBase):
- def __init__(
- self,
+class BlendFilterNoise(CustomNoiseItemBase):
+ def __init__(
+ self,
+ factor,
+ *,
+ noise,
+ blend_mode,
+ ffilter,
+ ffilter_scale,
+ ffilter_strength,
+ ffilter_threshold,
+ enhance_mode,
+ enhance_strength,
+ affect,
+ normalize_result,
+ normalize_noise,
+ ):
+ if len(noise.items) == 0:
+ raise ValueError("BlendFilterNoise requires at least one noise item")
+ super().__init__(
factor,
- *,
- noise,
- blend_mode,
- ffilter,
- ffilter_scale,
- ffilter_strength,
- ffilter_threshold,
- enhance_mode,
- enhance_strength,
- affect,
- normalize_result,
- normalize_noise,
- ):
- if len(noise.items) == 0:
- raise ValueError("BlendFilterNoise requires at least one noise item")
- super().__init__(
- factor,
- noise=noise.clone(),
- blend_mode=blend_mode,
- ffilter=ffilter,
- ffilter_scale=ffilter_scale,
- ffilter_strength=ffilter_strength,
- ffilter_threshold=ffilter_threshold,
- enhance_mode=enhance_mode,
- enhance_strength=enhance_strength,
- affect=affect,
- normalize_result=normalize_result,
- normalize_noise=normalize_noise,
+ noise=noise.clone(),
+ blend_mode=blend_mode,
+ ffilter=ffilter,
+ ffilter_scale=ffilter_scale,
+ ffilter_strength=ffilter_strength,
+ ffilter_threshold=ffilter_threshold,
+ enhance_mode=enhance_mode,
+ enhance_strength=enhance_strength,
+ affect=affect,
+ normalize_result=normalize_result,
+ normalize_noise=normalize_noise,
+ )
+
+ def clone_key(self, k):
+ if k == "noise":
+ return self.noise.clone()
+ return super().clone_key(k)
+
+ def apply_effects(self, noise, sigma):
+ blu = external.MODULES.bleh.py.latent_utils
+ if self.ffilter:
+ noise = blu.ffilter(
+ noise,
+ self.ffilter_threshold,
+ self.ffilter_scale,
+ self.ffilter,
+ self.ffilter_strength,
)
+ if self.enhance_mode != "none" and self.enhance_strength != 0:
+ noise = blu.enhance_tensor(
+ noise,
+ self.enhance_mode,
+ self.enhance_strength,
+ sigma=sigma,
+ skip_multiplier=0,
+ adjust_scale=False,
+ )
+ return noise
- def clone_key(self, k):
- if k == "noise":
- return self.noise.clone()
- return super().clone_key(k)
+ def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
+ factor = self.factor
+ noise_items = self.noise.items
+ noise_samplers = tuple(
+ ni.make_noise_sampler(x, *args, normalized=False, **kwargs)
+ for ni in noise_items
+ )
+ num_samplers = len(noise_samplers)
+ normalize_noise = self.get_normalize(
+ "normalize_noise",
+ normalized or num_samplers > 1,
+ )
+ normalize_result = self.get_normalize("normalize_result", normalized)
+ noise_effects = self.affect in {"noise", "both"}
+ result_effects = self.affect in {"result", "both"}
+ noise_init = torch.zeros_like(x)
- def apply_effects(self, noise, sigma):
- if self.ffilter:
- noise = BLU.ffilter(
- noise,
- self.ffilter_threshold,
- self.ffilter_scale,
- self.ffilter,
- self.ffilter_strength,
- )
- if self.enhance_mode != "none" and self.enhance_strength != 0:
- noise = BLU.enhance_tensor(
- noise,
- self.enhance_mode,
- self.enhance_strength,
- sigma=sigma,
- skip_multiplier=0,
- adjust_scale=False,
- )
+ def noise_sampler(s, sn):
+ noise = noise_init.clone()
+ for ni, ns in zip(noise_items, noise_samplers):
+ curr_noise = scale_noise(ns(s, sn), normalized=normalize_noise)
+ if noise_effects:
+ curr_noise = self.apply_effects(curr_noise, s)
+ if self.blend_mode == "simple_add":
+ noise += curr_noise.mul_(ni.factor)
+ else:
+ noise = utils.BLENDING_MODES[self.blend_mode](
+ noise,
+ curr_noise,
+ ni.factor,
+ )
+ del curr_noise
+ noise = scale_noise(noise, factor, normalized=normalize_result)
+ if result_effects:
+ noise = self.apply_effects(noise, s)
return noise
- def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
- factor = self.factor
- noise_items = self.noise.items
- noise_samplers = tuple(
- ni.make_noise_sampler(x, *args, normalized=False, **kwargs)
- for ni in noise_items
- )
- num_samplers = len(noise_samplers)
- normalize_noise = self.get_normalize(
- "normalize_noise",
- normalized or num_samplers > 1,
- )
- normalize_result = self.get_normalize("normalize_result", normalized)
- noise_effects = self.affect in {"noise", "both"}
- result_effects = self.affect in {"result", "both"}
- noise_init = torch.zeros_like(x)
+ return noise_sampler
- def noise_sampler(s, sn):
- noise = noise_init.clone()
- for ni, ns in zip(noise_items, noise_samplers):
- curr_noise = scale_noise(ns(s, sn), normalized=normalize_noise)
- if noise_effects:
- curr_noise = self.apply_effects(curr_noise, s)
- if self.blend_mode == "simple_add":
- noise += curr_noise.mul_(ni.factor)
- else:
- noise = BLU.BLENDING_MODES[self.blend_mode](
- noise,
- curr_noise,
- ni.factor,
- )
- del curr_noise
- noise = scale_noise(noise, factor, normalized=normalize_result)
- if result_effects:
- noise = self.apply_effects(noise, s)
- return noise
- return noise_sampler
-
- class BlehOpsNoise(CustomNoiseItemBase):
- def __init__(
- self,
+class BlehOpsNoise(CustomNoiseItemBase):
+ def __init__(
+ self,
+ factor,
+ *,
+ noise,
+ rules,
+ normalize,
+ ):
+ if len(noise.items) == 0:
+ raise ValueError("BlehOpsNoise requires at least one noise item")
+ super().__init__(
factor,
- *,
- noise,
- rules,
- normalize,
- ):
- if len(noise.items) == 0:
- raise ValueError("BlehOpsNoise requires at least one noise item")
- super().__init__(
- factor,
- noise=noise.clone(),
- rules=rules,
- normalize=normalize,
- )
+ noise=noise.clone(),
+ rules=rules,
+ normalize=normalize,
+ )
- def clone_key(self, k):
- if k == "noise":
- return self.noise.clone()
- return super().clone_key(k)
+ def clone_key(self, k):
+ if k == "noise":
+ return self.noise.clone()
+ return super().clone_key(k)
- def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
- factor = self.factor
- normalize = self.get_normalize("normalize", normalized)
- rulegroup = self.rules
- internal_ns = self.noise.make_noise_sampler(
- x,
- *args,
- normalized=False,
- **kwargs,
- )
+ def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
+ bops = external.MODULES.bleh.py.nodes.ops
+ factor = self.factor
+ normalize = self.get_normalize("normalize", normalized)
+ rulegroup = self.rules
+ internal_ns = self.noise.make_noise_sampler(
+ x,
+ *args,
+ normalized=False,
+ **kwargs,
+ )
- def noise_sampler(s, sn):
- noise = internal_ns(s, sn)
- if len(rulegroup.rules):
- state = {
- BOPS.CondType.TYPE: BOPS.PatchType.LATENT,
- BOPS.CondType.PERCENT: 0.0,
- BOPS.CondType.BLOCK: -1,
- BOPS.CondType.STAGE: -1,
- "sigma": None if s is None else s,
- "h": noise,
- "hsp": x.detach().clone(),
- "target": "h",
- }
- noise = rulegroup.eval(state, toplevel=True)["h"]
- return scale_noise(noise, factor, normalized=normalize)
+ def noise_sampler(s, sn):
+ noise = internal_ns(s, sn)
+ if len(rulegroup.rules):
+ state = {
+ bops.CondType.TYPE: bops.PatchType.LATENT,
+ bops.CondType.PERCENT: 0.0,
+ bops.CondType.BLOCK: -1,
+ bops.CondType.STAGE: -1,
+ "sigma": None if s is None else s,
+ "h": noise,
+ "hsp": x.detach().clone(),
+ "target": "h",
+ }
+ noise = rulegroup.eval(state, toplevel=True)["h"]
+ return scale_noise(noise, factor, normalized=normalize)
- return noise_sampler
+ return noise_sampler
NOISE_SAMPLERS: dict[NoiseType, Callable] = {
diff --git a/py/noise_generation.py b/py/noise_generation.py
index 98b08fb..2fe42e1 100644
--- a/py/noise_generation.py
+++ b/py/noise_generation.py
@@ -18,13 +18,8 @@ try:
except ImportError:
HAVE_WAVELETS = False
-from .noise_utils import (
- BLENDING_MODES,
- quantile_normalize,
- scale_noise,
- scale_samples,
- tensor_to,
-)
+from . import utils
+from .utils import quantile_normalize, scale_noise, tensor_to
# ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311
@@ -439,7 +434,7 @@ class PerlinOldNoiseGenerator(NoiseGenerator):
return cls.perlin_noise_tensor(vectors, positions, blend=blend).squeeze(0)
def generate(self, *_args):
- blend = BLENDING_MODES[self.blend_mode]
+ blend = utils.BLENDING_MODES[self.blend_mode]
noise = self.rand_like(fun=torch.rand).div_(self.div_fac)
_batch, channels, noise_height, noise_width = noise.shape
@@ -528,7 +523,7 @@ class HighresPyramidNoiseGenerator(NoiseGenerator):
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 += scale_samples(
+ noise += utils.scale_samples(
tensor_to(torch.randn(b, c, h, w, generator=self.generator), noise),
orig_w,
orig_h,
@@ -565,7 +560,7 @@ class PyramidOldNoiseGenerator(NoiseGenerator):
r = 1
for i in range(self.iterations):
r *= 2
- noise += scale_samples(
+ noise += utils.scale_samples(
torch.normal(
mean=0,
std=0.5**i,
@@ -608,7 +603,7 @@ class PyramidNoiseGenerator(NoiseGenerator):
torch.rand(1, generator=self.generator).cpu().item() * 2 + 2
) # Rather than always going 2x,
w, h = max(1, int(w / (r**i))), max(1, int(h / (r**i)))
- noise += scale_samples(
+ noise += utils.scale_samples(
torch.randn(
b,
c,
diff --git a/py/powernoise.py b/py/powernoise.py
index 914f66a..bfff7bb 100644
--- a/py/powernoise.py
+++ b/py/powernoise.py
@@ -23,7 +23,7 @@ from .nodes import (
SonarNormalizeNoiseNodeMixin,
)
from .noise import CustomNoiseItemBase
-from .noise_utils import scale_noise
+from .utils import scale_noise
# ruff: noqa: ANN003, FBT001, FBT002
diff --git a/py/sonar.py b/py/sonar.py
index fc0f8fe..e30bad0 100644
--- a/py/sonar.py
+++ b/py/sonar.py
@@ -9,19 +9,19 @@ from sys import stderr
from typing import Any, Callable, NamedTuple
import torch
-from comfy.k_diffusion import sampling
+from comfy.k_diffusion.sampling import get_ancestral_step, to_d
from comfy.samplers import KSampler, k_diffusion_sampling
from torch import Tensor
from tqdm.auto import trange
-from . import noise
-from .noise_utils import BLENDING_MODES
+from . import noise, utils
class HistoryType(Enum):
ZERO = auto()
RAND = auto()
SAMPLE = auto()
+ SAMPLE_NORM = auto()
class GuidanceType(Enum):
@@ -37,34 +37,106 @@ class GuidanceConfig(NamedTuple):
latent: Tensor | None = None
+class MomentumMode(Enum):
+ CLASSIC = auto()
+ NEW = auto()
+ DENOISED = auto()
+
+
class SonarConfig(NamedTuple):
momentum: float = 0.95
momentum_hist: float = 0.75
direction: float = 1.0
+ momentum_start_step: int = 0
+ momentum_end_step: int = 9999
+ always_update_history: bool = True
+ momentum_mode: MomentumMode = MomentumMode.NEW
init: HistoryType = HistoryType.ZERO
noise_type: noise.NoiseType | None = None
custom_noise: noise.CustomNoise | None = None
rand_init_noise_type: noise.NoiseType | None = None
+ rand_init_noise_multiplier: float | int = 1.0
guidance: GuidanceConfig | None = None
blend_mode: str = "lerp"
+ momentum_blend_mode: str | None = None
+ history_blend_mode: str | None = None
+ guidance_blend_mode: str | None = None
+
+ def get_with_default(self, k: str, default: Any) -> Any: # noqa: ANN401
+ val = getattr(self, k)
+ return val if val is not None else default
class SonarBase:
DEFAULT_NOISE_TYPE = noise.NoiseType.GAUSSIAN
- def __init__(self, cfg: SonarConfig, *, blend_mode="lerp") -> None:
+ def __init__(self, cfg: SonarConfig) -> None:
self.history_d = None
self.cfg = cfg
self.noise_sampler = None
- self.blend_function = BLENDING_MODES[blend_mode]
+ blend_mode = cfg.blend_mode
+ momentum_blend_mode = cfg.get_with_default("momentum_blend_mode", blend_mode)
+ history_blend_mode = cfg.get_with_default("history_blend_mode", blend_mode)
+ guidance_blend_mode = cfg.get_with_default("guidance_blend_mode", blend_mode)
+ bf = self.blend = utils.BLENDING_MODES[blend_mode]
+ self.momentum_blend = (
+ bf
+ if momentum_blend_mode == blend_mode
+ else utils.BLENDING_MODES[momentum_blend_mode]
+ )
+ self.history_blend = (
+ bf
+ if history_blend_mode == blend_mode
+ else utils.BLENDING_MODES[history_blend_mode]
+ )
+ self.guidance_blend = (
+ bf
+ if guidance_blend_mode == blend_mode
+ else utils.BLENDING_MODES[guidance_blend_mode]
+ )
+
+ _cfg_fixups = (
+ ("momentum_mode", MomentumMode),
+ ("init", HistoryType),
+ ("noise_type", noise.NoiseType),
+ )
+
+ @classmethod
+ def get_config(
+ cls,
+ cfg: SonarConfig | None = None,
+ ext: dict | None = None,
+ ) -> SonarConfig:
+ cfgdict = ext.copy() if ext is not None else {}
+ empty = object()
+ for k, enum_class in cls._cfg_fixups:
+ val = cfgdict.get(k, empty)
+ if val is empty:
+ continue
+ if isinstance(val, str):
+ val = getattr(enum_class, val.strip().upper(), empty)
+ if val is empty:
+ validstr = ", ".join(enum_class.__members__.keys())
+ errstr = f"Bad value for {k} of type enum {enum_class.__name__}, must be one of the following: {validstr}"
+ raise ValueError(errstr)
+ cfgdict[k] = val
+ continue
+ if not isinstance(val, enum_class):
+ errstr = f"Bad parameter type for {k}: Must be valid string or instance of {enum_class.__name__}"
+ raise TypeError(errstr)
+
+ if cfg is None:
+ return SonarConfig(**cfgdict)
+ cfgdict = cfg._asdict() | cfgdict
+ return SonarConfig(**cfgdict)
def set_noise_sampler(
self,
x: Tensor,
- sigmas,
+ sigmas: Tensor,
noise_sampler: Callable | None,
seed: int | None = None,
- ):
+ ) -> Callable:
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
if noise_sampler is not None and self.cfg.noise_type not in {
None,
@@ -94,17 +166,32 @@ class SonarBase:
self.noise_sampler = noise_sampler
return noise_sampler
- def init_hist_d(self, x: Tensor) -> None:
- if self.history_d is not None:
+ def init_hist_d(
+ self,
+ x: Tensor,
+ denoised: Tensor,
+ sigma: Tensor,
+ *,
+ step: int,
+ ) -> None:
+ if self.history_d is not None or not self.check_step(step, is_history=True):
return
+ cfg = self.cfg
+ init = cfg.init
# memorize delta momentum
- if self.cfg.init == HistoryType.ZERO:
+ if init == HistoryType.ZERO:
self.history_d = None
- elif self.cfg.init == HistoryType.SAMPLE:
- self.history_d = x
- elif self.cfg.init == HistoryType.RAND:
+ elif init == HistoryType.SAMPLE:
+ self.history_d = (
+ x if cfg.momentum_mode != MomentumMode.DENOISED else denoised
+ )
+ elif init == HistoryType.SAMPLE_NORM:
+ self.history_d = (
+ x if cfg.momentum_mode != MomentumMode.DENOISED else denoised
+ ) / sigma
+ elif init == HistoryType.RAND:
ns = noise.get_noise_sampler(
- self.cfg.rand_init_noise_type,
+ cfg.rand_init_noise_type,
x,
None,
None,
@@ -113,6 +200,8 @@ class SonarBase:
normalized=True,
)
self.history_d = ns(None, None)
+ if cfg.rand_init_noise_multiplier != 1:
+ self.history_d *= cfg.rand_init_noise_multiplier
else:
raise ValueError("Sonar sampler: bad history type")
@@ -129,31 +218,106 @@ class SonarBase:
direction,
)
- def update_hist(self, momentum_d: torch.Tensor) -> None:
+ def check_step(self, step: int, *, is_history: bool = False):
+ cfg = self.cfg
+ if is_history and cfg.always_update_history:
+ return True
+ return cfg.momentum_start_step <= step <= cfg.momentum_end_step
+
+ def update_hist(self, momentum_d: torch.Tensor, step: int) -> None:
hd, cfg = self.history_d, self.cfg
- if cfg.momentum_hist == 1:
+ if cfg.momentum_hist == 1 or not self.check_step(step, is_history=True):
return
hd_ratio, hd_scale, md_scale = self.history_ratios
self.history_d = (
momentum_d
if hd is None
- else self.blend_function(momentum_d * md_scale, hd * hd_scale, hd_ratio)
+ else self.history_blend(momentum_d * md_scale, hd * hd_scale, hd_ratio)
)
- def momentum_step(self, x: Tensor, d: Tensor, dt: Tensor):
- momentum = self.cfg.momentum
- if momentum == 1.0:
- return x + d * dt
+ def momentum_mix(
+ self,
+ history: Tensor | None,
+ item: Tensor,
+ sigma: Tensor,
+ *,
+ is_denoised: bool = False,
+ momentum=None,
+ ) -> Tensor:
+ momentum = self.cfg.momentum if momentum is None else momentum
+ mode = self.cfg.momentum_mode
+ if (
+ momentum == 1 # noqa: PLR0916
+ or history is None
+ or (mode == MomentumMode.DENOISED and not is_denoised)
+ or (mode != MomentumMode.DENOISED and is_denoised)
+ ):
+ return item
+ return self.momentum_blend(
+ history * sigma if is_denoised else history,
+ item,
+ momentum,
+ )
+
+ def get_momentum_denoised(
+ self,
+ x: Tensor,
+ denoised: Tensor,
+ sigma: Tensor,
+ *,
+ step: int,
+ momentum: float | None = None,
+ update_history=True,
+ ) -> Tensor:
hd = self.history_d
- # correct current `d` with momentum
- momentum_d = d if hd is None else self.blend_function(hd, d, momentum)
+ momentum_denoised = self.momentum_mix(
+ hd,
+ denoised,
+ sigma,
+ is_denoised=True,
+ momentum=momentum,
+ )
+ if update_history:
+ self.init_hist_d(x, denoised, sigma, step=step)
+ self.update_hist(denoised / sigma, step=step)
+ return momentum_denoised if self.check_step(step) else denoised
- # Euler method with momentum
- x = x + momentum_d * dt # noqa: PLR6104
+ def get_momentum_d(
+ self,
+ x: Tensor,
+ denoised: Tensor,
+ sigma: Tensor,
+ *,
+ step: int,
+ momentum: float | None = None,
+ d: Tensor | None = None,
+ update_history=True,
+ ) -> Tensor:
+ hd = self.history_d
+ cfg = self.cfg
+ momentum = cfg.momentum if momentum is None else momentum
+ mode = cfg.momentum_mode
+ d = to_d(x, sigma, denoised) if d is None else d
+ if momentum == 1 or mode == MomentumMode.DENOISED:
+ return d
+ momentum_d = self.momentum_mix(hd, d, sigma)
+ if update_history:
+ self.init_hist_d(x, denoised, sigma, step=step)
+ self.update_hist(d if mode == MomentumMode.NEW else momentum_d, step=step)
+ return momentum_d if self.check_step(step) else d
- self.update_hist(momentum_d)
-
- return x
+ def momentum_step(
+ self,
+ step: int,
+ x: Tensor,
+ denoised: Tensor,
+ sigma: Tensor,
+ sigma_down: Tensor,
+ ) -> Tensor:
+ dt = sigma_down - sigma
+ denoised = self.get_momentum_denoised(x, denoised, sigma, step=step)
+ momentum_d = self.get_momentum_d(x, denoised, sigma, step=step)
+ return (momentum_d * dt).add_(x)
class SonarGuidanceMixin:
@@ -174,9 +338,9 @@ class SonarGuidanceMixin:
return None
avg_s = latent.mean(dim=(-2, -1), keepdim=True)
std_s = latent.std(dim=(-2, -1), keepdim=True)
- return ((latent - avg_s) / std_s).to(latent.dtype)
+ return (latent - avg_s).div_(std_s).to(latent.dtype)
- def guidance_step(self, step_index: int, x: Tensor, denoised: Tensor):
+ def guidance_step(self, step_index: int, x: Tensor, denoised: Tensor) -> Tensor:
if (
self.guidance is None
or self.guidance.factor == 0.0
@@ -186,7 +350,12 @@ class SonarGuidanceMixin:
if self.ref_latent.device != x.device:
self.ref_latent = self.ref_latent.to(device=x.device)
if self.guidance.guidance_type == GuidanceType.LINEAR:
- return self.guidance_linear(x, self.ref_latent, self.guidance.factor)
+ return self.guidance_linear(
+ x,
+ self.ref_latent,
+ self.guidance.factor,
+ blend=self.guidance_blend,
+ )
if self.guidance.guidance_type == GuidanceType.EULER:
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
return self.guidance_euler(
@@ -212,16 +381,22 @@ class SonarGuidanceMixin:
std_t = denoised.std(dim=(-3, -2, -1), keepdim=True)
ref_img_shift = ref_latent * std_t + avg_t
- d = sampling.to_d(x, sigma, ref_img_shift)
+ d = to_d(x, sigma, ref_img_shift)
dt = (sigma_next - sigma) * factor
- return x + d * dt
+ return (d * dt).add_(x)
@staticmethod
- def guidance_linear(x: Tensor, ref_latent: Tensor, factor: float = 0.2) -> Tensor:
+ def guidance_linear(
+ x: Tensor,
+ ref_latent: Tensor,
+ factor: float = 0.2,
+ *,
+ blend=torch.lerp,
+ ) -> Tensor:
avg_t = x.mean(dim=(-3, -2, -1), keepdim=True)
std_t = x.std(dim=(-3, -2, -1), keepdim=True)
- ref_img_shift = ref_latent * std_t + avg_t
- return (1.0 - factor) * x + factor * ref_img_shift
+ ref_img_shift = (ref_latent * std_t).add_(avg_t)
+ return blend(x, ref_img_shift, factor)
class SonarWithGuidance(SonarBase, SonarGuidanceMixin):
@@ -234,9 +409,9 @@ class SonarSampler(SonarWithGuidance):
def __init__(
self,
model,
- sigmas,
- s_in,
- extra_args,
+ sigmas: Tensor,
+ s_in: Tensor,
+ extra_args: dict[str, Any],
*args: list[Any],
**kwargs: dict[str, Any],
):
@@ -246,6 +421,21 @@ class SonarSampler(SonarWithGuidance):
self.s_in = s_in
self.extra_args = extra_args
+ def call_model(
+ self,
+ x: Tensor,
+ sigma: Tensor,
+ *args: list[Any],
+ s_in=None,
+ extra_args=None,
+ ) -> Tensor:
+ if s_in is None:
+ s_in = self.s_in
+ extra_args = (
+ self.extra_args if extra_args is None else self.extra_args | extra_args
+ )
+ return self.model(x, sigma * s_in, *args, **extra_args)
+
class SonarEuler(SonarSampler):
def __init__(
@@ -256,16 +446,18 @@ class SonarEuler(SonarSampler):
super().__init__(*args, **kwargs)
def step(self, step_index: int, sample: torch.FloatTensor):
- self.init_hist_d(sample)
- sigma, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1]
+ sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
- denoised = self.model(sample, sigma * self.s_in, **self.extra_args)
- derivative = sampling.to_d(sample, sigma, denoised)
- dt = sigma_to - sigma
+ denoised = self.call_model(sample, sigma)
+ result_sample = self.momentum_step(
+ step_index,
+ sample,
+ denoised,
+ sigma,
+ sigma_next,
+ )
- result_sample = self.momentum_step(sample, derivative, dt)
-
- if sigma_to > 0:
+ if sigma_next > 0:
result_sample = self.guidance_step(step_index, result_sample, denoised)
return (
@@ -280,17 +472,16 @@ class SonarEuler(SonarSampler):
def sampler(
cls,
model,
- x,
- sigmas,
- extra_args=None,
+ x: Tensor,
+ sigmas: Tensor,
+ extra_args: dict | None = None,
callback=None,
- disable=None,
+ disable: bool | None = None, # noqa: FBT001
noise_sampler: Callable | None = None,
- sonar_blend_mode="lerp",
- sonar_config=None,
- ):
- if sonar_config is None:
- sonar_config = SonarConfig()
+ sonar_config: SonarConfig | None = None,
+ sonar_params: dict | None = None,
+ ) -> Tensor:
+ sonar_config = cls.get_config(sonar_config, sonar_params)
s_in = x.new_ones((x.shape[0],))
sonar = cls(
model,
@@ -298,7 +489,6 @@ class SonarEuler(SonarSampler):
s_in,
{} if extra_args is None else extra_args,
sonar_config,
- blend_mode=sonar_blend_mode,
)
sonar.set_noise_sampler(
x,
@@ -342,31 +532,32 @@ class SonarEulerAncestral(SonarSampler):
step_index: int,
sample: torch.FloatTensor,
):
- self.init_hist_d(sample)
-
- sigma_from, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1]
- sigma_down, sigma_up = sampling.get_ancestral_step(
- sigma_from,
- sigma_to,
+ sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
+ sigma_down, sigma_up = get_ancestral_step(
+ sigma,
+ sigma_next,
eta=self.eta,
)
- denoised = self.model(sample, sigma_from * self.s_in, **self.extra_args)
- derivative = sampling.to_d(sample, sigma_from, denoised)
- dt = sigma_down - sigma_from
-
- result_sample = self.momentum_step(sample, derivative, dt)
- if sigma_to > 0:
+ denoised = self.call_model(sample, sigma)
+ result_sample = self.momentum_step(
+ step_index,
+ sample,
+ denoised,
+ sigma,
+ sigma_down,
+ )
+ if sigma_next > 0:
result_sample = self.guidance_step(step_index, result_sample, denoised)
result_sample = ( # noqa: PLR6104
result_sample
- + self.noise_sampler(sigma_from, sigma_to) * self.s_noise * sigma_up
+ + self.noise_sampler(sigma, sigma_next) * (self.s_noise * sigma_up)
)
return (
result_sample,
- sigma_from,
- sigma_from,
+ sigma,
+ sigma,
denoised,
)
@@ -380,15 +571,14 @@ class SonarEulerAncestral(SonarSampler):
extra_args=None,
callback=None,
disable=None,
- sonar_blend_mode="lerp",
- sonar_config=None,
+ sonar_config: SonarConfig | None = None,
+ sonar_params: dict | None = None,
eta=1.0,
s_noise=1.0,
noise_sampler: Callable | None = None,
):
- if sonar_config is None:
- sonar_config = SonarConfig()
- s_in = x.new_ones([x.shape[0]])
+ sonar_config = cls.get_config(sonar_config, sonar_params)
+ s_in = x.new_ones((x.shape[0],))
sonar = cls(
eta,
s_noise,
@@ -397,7 +587,6 @@ class SonarEulerAncestral(SonarSampler):
s_in,
{} if extra_args is None else extra_args,
sonar_config,
- blend_mode=sonar_blend_mode,
)
sonar.set_noise_sampler(
x,
@@ -439,114 +628,134 @@ class SonarDPMPPSDE(SonarSampler):
self.s_noise = s_noise
@staticmethod
- def sigma_fn(t) -> float:
+ def sigma_fn(t: Tensor) -> float:
return t.neg().exp()
@staticmethod
- def t_fn(sigma) -> float:
- return sigma.log.neg()
+ def t_fn(sigma: Tensor) -> float:
+ return sigma.log().neg()
# DPM++ solver algorithm copied from ComfyUI source.
def momentum_step( # noqa: PLR0914
self,
- step_index,
+ step_index: int,
x: Tensor,
denoised: Tensor,
- sigma_from,
- sigma_to,
- sigma_down,
- ):
- if sigma_to == 0:
- derivative = sampling.to_d(x, sigma_from, denoised)
- dt = sigma_down - sigma_from
- return super().momentum_step(x, derivative, dt)
+ sigma: Tensor,
+ sigma_next: Tensor,
+ sigma_down: Tensor,
+ ) -> Tensor:
+ if sigma_next == 0:
+ return super().momentum_step(step_index, x, denoised, sigma, sigma_down)
- def sigma_fn(t):
- return t.neg().exp()
-
- def t_fn(sigma):
- return sigma.log().neg()
-
- hd = self.history_d
+ cfg = self.cfg
# Halve the momentum proportion if there's history since we will use it twice.
adjusted_momentum = (
- self.cfg.momentum + (1 - self.cfg.momentum) / 2
- if hd is not None
- else self.cfg.momentum
+ cfg.momentum + (1 - cfg.momentum) / 2
+ if self.history_d is not None
+ else cfg.momentum
)
r = 1 / 2
# DPM-Solver++
- t, t_next = t_fn(sigma_from), t_fn(sigma_to)
+ t, t_next = self.t_fn(sigma), self.t_fn(sigma_next)
h = t_next - t
s = t + h * r
fac = 1 / (2 * r)
# Step 1
- sd, su = sampling.get_ancestral_step(sigma_fn(t), sigma_fn(s), self.eta)
- s_ = t_fn(sd)
- diff_2 = (t - s_).expm1() * denoised
- momentum_d = (
- diff_2 if hd is None else self.blend_function(hd, diff_2, adjusted_momentum)
- )
- self.update_hist(momentum_d)
- hd = self.history_d
- x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - momentum_d
- x_2 += self.noise_sampler(sigma_fn(t), sigma_fn(s)) * self.s_noise * su
- denoised_2 = self.model(x_2, sigma_fn(s) * self.s_in, **self.extra_args)
-
- # Step 2
- sd, su = sampling.get_ancestral_step(
- sigma_fn(t),
- sigma_fn(t_next),
+ s_t, s_s = self.sigma_fn(t), self.sigma_fn(s)
+ sd, su = get_ancestral_step(
+ s_t,
+ s_s,
self.eta,
)
- t_next_ = t_fn(sd)
- denoised_d = (1 - fac) * denoised + fac * denoised_2
- diff_1 = (t - t_next_).expm1() * denoised_d
- hd = self.history_d
- momentum_d = (
- diff_1 if hd is None else self.blend_function(hd, diff_1, adjusted_momentum)
+ s_ = self.t_fn(sd)
+ momentum_denoised = self.get_momentum_denoised(
+ x,
+ denoised,
+ sigma,
+ step=step_index,
)
- self.update_hist(momentum_d)
- x = (sigma_fn(t_next_) / sigma_fn(t)) * x - momentum_d
+ diff_2 = (t - s_).expm1() * momentum_denoised
+ momentum_d = self.get_momentum_d(
+ x,
+ momentum_denoised,
+ sigma,
+ step=step_index,
+ momentum=adjusted_momentum,
+ d=diff_2,
+ )
+ x_2 = ((self.sigma_fn(s_) / s_t) * x).sub_(momentum_d)
+ x_2 += self.noise_sampler(s_t, s_s).mul_(
+ self.s_noise * su,
+ )
+ sigma_2 = s_s
+ denoised_2 = self.call_model(x_2, sigma_2)
+ momentum_denoised_2 = self.get_momentum_denoised(
+ x,
+ denoised_2,
+ sigma_2,
+ step=step_index,
+ )
+
+ # Step 2
+ s_t_next = self.sigma_fn(t_next)
+ sd, su = get_ancestral_step(
+ s_t,
+ s_t_next,
+ self.eta,
+ )
+ t_down = self.t_fn(sd)
+ denoised_d = (1 - fac) * momentum_denoised + fac * momentum_denoised_2
+ diff_1 = (t - t_down).expm1() * denoised_d
+ momentum_d = self.get_momentum_d(
+ x,
+ momentum_denoised_2,
+ sigma_2,
+ step=step_index,
+ momentum=adjusted_momentum,
+ d=diff_1,
+ )
+ x = ((self.sigma_fn(t_down) / s_t) * x).sub_(momentum_d)
x = self.guidance_step(step_index, x, denoised_d)
- return x + self.noise_sampler(sigma_fn(t), sigma_fn(t_next)) * self.s_noise * su
+ x += self.noise_sampler(s_t, s_t_next).mul_(
+ self.s_noise * su,
+ )
+ return x
def step(
self,
step_index: int,
sample: torch.FloatTensor,
- ):
+ ) -> Tensor:
def sigma_fn(t):
return t.neg().exp()
def t_fn(sigma):
return sigma.log().neg()
- self.init_hist_d(sample)
-
- sigma_from, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1]
- sigma_down, _sigma_up = sampling.get_ancestral_step(
- sigma_from,
- sigma_to,
+ sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
+ sigma_down, _sigma_up = get_ancestral_step(
+ sigma,
+ sigma_next,
eta=self.eta,
)
- denoised = self.model(sample, sigma_from * self.s_in, **self.extra_args)
+ denoised = self.call_model(sample, sigma)
result_sample = self.momentum_step(
step_index,
sample,
denoised,
- sigma_from,
- sigma_to,
+ sigma,
+ sigma_next,
sigma_down,
)
return (
result_sample,
- sigma_from,
- sigma_from,
+ sigma,
+ sigma,
denoised,
)
@@ -555,20 +764,19 @@ class SonarDPMPPSDE(SonarSampler):
def sampler(
cls,
model,
- x,
- sigmas,
- extra_args=None,
+ x: Tensor,
+ sigmas: Tensor,
+ extra_args: dict | None = None,
callback=None,
- disable=None,
- sonar_blend_mode="lerp",
- sonar_config=None,
+ disable: bool | None = None, # noqa: FBT001
+ sonar_config: SonarConfig | None = None,
+ sonar_params: dict | None = None,
eta=1.0,
s_noise=1.0,
noise_sampler=None,
- ):
- if sonar_config is None:
- sonar_config = SonarConfig()
- s_in = x.new_ones([x.shape[0]])
+ ) -> Tensor:
+ sonar_config = cls.get_config(sonar_config, sonar_params)
+ s_in = x.new_ones((x.shape[0],))
sonar = cls(
eta,
s_noise,
@@ -577,7 +785,6 @@ class SonarDPMPPSDE(SonarSampler):
s_in,
{} if extra_args is None else extra_args,
sonar_config,
- blend_mode=sonar_blend_mode,
)
sonar.set_noise_sampler(
x,
@@ -604,7 +811,7 @@ class SonarDPMPPSDE(SonarSampler):
return x
-def add_samplers():
+def add_samplers() -> None:
extra_samplers = {
"sonar_euler": SonarEuler.sampler,
"sonar_euler_ancestral": SonarEulerAncestral.sampler,
diff --git a/py/noise_utils.py b/py/utils.py
similarity index 86%
rename from py/noise_utils.py
rename to py/utils.py
index 5d4366e..c0cf222 100644
--- a/py/noise_utils.py
+++ b/py/utils.py
@@ -8,6 +8,41 @@ from comfy.utils import common_upscale
from .external import MODULES as EXT
+BLENDING_MODES = {"lerp": torch.lerp}
+UPSCALE_METHODS = (
+ "bilinear",
+ "nearest-exact",
+ "nearest",
+ "area",
+ "bicubic",
+ "bislerp",
+)
+
+
+def scale_samples(
+ samples: torch.Tensor,
+ width: int,
+ height: int,
+ *,
+ mode: str = "bicubic",
+) -> torch.Tensor:
+ return common_upscale(samples, width, height, mode, None)
+
+
+def init_integrations(integrations) -> None:
+ global scale_samples, BLENDING_MODES, UPSCALE_METHODS # noqa: PLW0603
+
+ bleh = integrations.bleh
+ if bleh is None:
+ return
+ bleh_latentutils = bleh.py.latent_utils
+ BLENDING_MODES = bleh_latentutils.BLENDING_MODES
+ UPSCALE_METHODS = bleh_latentutils.UPSCALE_METHODS
+ scale_samples = bleh_latentutils.scale_samples
+
+
+EXT.register_init_handler(init_integrations)
+
def scale_noise(
noise: torch.Tensor,
@@ -33,22 +68,6 @@ def scale_noise(
return noise.mul_(factor) if factor != 1 else noise
-if "bleh" in EXT:
- scale_samples = EXT["bleh"].py.latent_utils.scale_samples
- BLENDING_MODES = EXT["bleh"].py.latent_utils.BLENDING_MODES
-else:
- BLENDING_MODES = {"lerp": torch.lerp}
-
- def scale_samples(
- samples: torch.Tensor,
- width: int,
- height: int,
- *,
- mode: str = "bicubic",
- ) -> torch.Tensor:
- return common_upscale(samples, width, height, mode, None)
-
-
CAN_NONBLOCK = {}
diff --git a/ruff.toml b/ruff.toml
index 3c32a79..e589ee0 100644
--- a/ruff.toml
+++ b/ruff.toml
@@ -15,6 +15,8 @@ ignore = [
"D102",
"D103",
"D104",
+ "D105",
+ "D106",
"D107",
"D211",
"D213",