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",