diff --git a/__init__.py b/__init__.py index 623f457..24691a8 100644 --- a/__init__.py +++ b/__init__.py @@ -11,5 +11,7 @@ NODE_CLASS_MAPPINGS = { "OCS SimpleRestartSchedule": nodes.SimpleRestartSchedule, "OCS ApplyFilterLatent": nodes.ApplyFilterLatent, "OCS ApplyFilterImage": nodes.ApplyFilterImage, + "OCS ExpressionFilteredLatentOperation": nodes.ExpressionFilteredLatentOperationNode, + "OCS ExpressionFilteredModelPatch": nodes.ExpressionFilteredModelPatchNode, } | custom_noise.NODE_CLASS_MAPPINGS __all__ = ["NODE_CLASS_MAPPINGS"] diff --git a/py/custom_noise/base.py b/py/custom_noise/base.py index 78ee4c6..2349535 100644 --- a/py/custom_noise/base.py +++ b/py/custom_noise/base.py @@ -1,11 +1,11 @@ import abc +from typing import Any, Callable + import torch -from typing import Callable, Any - from ..external import IntegratedNode +from ..nodes import NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE from ..noise import scale_noise -from ..nodes import WILDCARD_NOISE, NOISE_INPUT_TYPES_HINT class CustomNoiseItemBase(abc.ABC): diff --git a/py/custom_noise/nodes.py b/py/custom_noise/nodes.py index e04239b..670fb07 100644 --- a/py/custom_noise/nodes.py +++ b/py/custom_noise/nodes.py @@ -1,25 +1,28 @@ +from __future__ import annotations + import functools import inspect import math +from collections.abc import Callable +from typing import Any, NamedTuple + +import comfy.samplers import torch import yaml - -from typing import NamedTuple - from comfy.model_management import get_torch_device from comfy.model_patcher import set_model_options_post_cfg_function -import comfy.samplers +from tqdm import tqdm -from .base import CustomNoiseNodeBase, NormalizeNoiseNodeMixin, CustomNoiseItemBase -from .noise_perlin import DEFAULTS as PERLIN_DEFAULTS -from .noise_perlin import PerlinItem -from .noise_immiscibleref import ImmiscibleReferenceItem - -from ..external import MODULES, IntegratedNode from .. import filtering -from ..nodes import WILDCARD_NOISE, NOISE_INPUT_TYPES_HINT -from ..utils import scale_noise +from ..external import MODULES, IntegratedNode +from ..nodes import NOISE_INPUT_TYPES_HINT, WILDCARD_NOISE from ..noise import ImmiscibleNoise +from ..utils import scale_noise +from .base import CustomNoiseItemBase, CustomNoiseNodeBase, NormalizeNoiseNodeMixin +from .noise_immiscibleref import ImmiscibleReferenceItem +from .noise_perlin import Perlin, PerlinItem + +PERLIN_DEFAULTS = Perlin() class ToSonarNode: @@ -54,6 +57,7 @@ class PerlinAdvancedNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): def INPUT_TYPES(cls): MODULES.initialize() result = super().INPUT_TYPES() + blend_modes = tuple(filtering.BLENDING_MODES.keys()) result["required"] |= { "depth": ( "INT", @@ -154,6 +158,8 @@ class PerlinAdvancedNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): "INT", { "default": PERLIN_DEFAULTS.max_depth, + "min": -99999, + "max": 99999, "tooltip": "Basically crops the depth dimension to the specified value (inclusive). Negative values start from the end, the default of -1 does no cropping. Only has an effect when depth is non-zero.", }, ), @@ -179,16 +185,16 @@ class PerlinAdvancedNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): }, ), "blend": ( - tuple(filtering.BLENDING_MODES.keys()), + blend_modes, { - "default": "lerp", + "default": PERLIN_DEFAULTS.blend.name, "tooltip": "Blending function used when generating Perlin noise. When set to values other than LERP may not work at all or may not actually generate Perlin noise.", }, ), "pattern_break_blend": ( - tuple(filtering.BLENDING_MODES.keys()), + blend_modes, { - "default": "lerp", + "default": PERLIN_DEFAULTS.pattern_break_blend.name, "tooltip": "Blending function used to blend pattern broken noise with raw noise.", }, ), @@ -273,6 +279,74 @@ class PerlinAdvancedNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): }, ), } + opts = result.get("optional", {}) + opts |= { + "ridge_weight": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.ridge_weight, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Blend strength for blending in the ridge-adjusted noise.", + }, + ), + "ridge_scale": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.ridge_scale, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Controls the amplitude of generated ridges.", + }, + ), + "ridge_blend": ( + blend_modes, + { + "default": PERLIN_DEFAULTS.ridge_blend.name, + "tooltip": "Blend mode used for blending in the ridge-adjusted noise.", + }, + ), + "warp_strength": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.warp_strength, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Strength of domain warping. Shifts coordinates of later octaves using the outputs from previous octaves.", + }, + ), + "octave_shift": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.octave_shift, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Shifts gradient dimensions (height, width, depth if enabled) across octaves. Shift is multiplied by the 1-based octave. I.E. shift 0.75 at the first octave is 1*0.75 and this is rounded to the nearest integer value. You can use this to only shift every other octave, etc. The shift can be negative to roll in the other direction.", + }, + ), + "curl_strength": ( + "FLOAT", + { + "default": PERLIN_DEFAULTS.curl_strength, + "min": -10000.0, + "max": 10000.0, + "tooltip": "Twists the gradients with additional Gaussian noise.", + }, + ), + "curl_dims": ( + "STRING", + { + "default": PERLIN_DEFAULTS.get_commasep("curl_dims"), + "tooltip": "Comma-separated axes to apply the curl effect to. Only does something when curl_strength is non-zero. In 3D mode the axes would be 0 (depth), 1 (height), 2 (width). In 2D mode you'd have 0 (height) and 1 (depth).", + }, + ), + "base_noise_opt": ( + WILDCARD_NOISE, + { + "tooltip": f"Optional input for noise to use as the Perlin base.\n{NOISE_INPUT_TYPES_HINT}", + }, + ), + } return result @classmethod @@ -297,10 +371,16 @@ class PerlinSimpleNode(PerlinAdvancedNode): def INPUT_TYPES(cls): result = super().INPUT_TYPES() orig_reqs = result["required"] + orig_opts = result["optional"] reqs = {k: v for k, v in orig_reqs.items() if k in cls._COPY_KEYS} reqs["lacunarity"] = orig_reqs["lacunarity_height"] reqs["res"] = orig_reqs["res_height"] result["required"] = reqs + result["optional"] = ( + {"ocs_noise_opt": orig_opts["ocs_noise_opt"]} + if "ocs_noise_opt" in orig_opts + else {} + ) return result @classmethod @@ -881,21 +961,34 @@ class ImmiscibleConfig(NamedTuple): maximize: bool = False distance_scale: float = 0.1 distance_scale_ref: float = 0.1 + use_triton: bool = True + abs_mode: bool = False + abs_distance_mode: bool = False + shuffle_seed: int | None = None start_time: float = 0.0 end_time: float = 1.0 blend: float = 1.0 blend_mode: str = "lerp" + force: bool = False + operation_reference: Callable | None = None + operation_noise: Callable | None = None + operation_result: Callable | None = None + filter_reference: filtering.Filter | None = None + filter_noise: filtering.Filter | None = None + filter_result: filtering.Filter | None = None + filter_postcfg: filtering.Filter | None = None DEFAULT_IMMISCIBLE_CONFIG = ImmiscibleConfig() class OverrideSamplerConfig(NamedTuple): + verbose: bool = False sampler: object | None = None sampler_kwargs: dict | None = None noise_start_time: float = 0.0 noise_end_time: float = 1.0 - cpu_noise: bool = True + cpu_noise: bool = False normalize: bool = True force_params: list | tuple = () custom_noise: object | None = None @@ -962,8 +1055,11 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode): "noise_prediction", "uncond_sub_cond", "cond_sub_uncond", + "denoised_sub_uncond", "latent", "noise", + "initial_latent", + "noise_gen_prev", ), { "default": DICFG.reference, @@ -1115,6 +1211,9 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode): "tooltip": "Optional input for blended noise (only used when blend is not 1.0). Can be used if you want to blend with a different noise type.", }, ), + "operation_reference": ("LATENT_OPERATION",), + "operation_noise": ("LATENT_OPERATION",), + "operation_result": ("LATENT_OPERATION",), }, } @@ -1142,7 +1241,10 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode): custom_noise_opt=None, custom_noise_ref=None, custom_noise_blend=None, - latent_reference: dict | None = None, + latent_ref: dict | None = None, + operation_reference: Callable | None = None, + operation_noise: Callable | None = None, + operation_result: Callable | None = None, ): MODULES.initialize() sampler_kwargs = {} @@ -1165,7 +1267,11 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode): "end_time": immiscible_end_time, "blend": immiscible_blend, "blend_mode": immiscible_blend_mode, + "operation_reference": operation_reference, + "operation_noise": operation_noise, + "operation_result": operation_result, } + ocs_verbose = False if yaml_parameters: extra_params = yaml.safe_load(yaml_parameters) if extra_params is None: @@ -1176,19 +1282,29 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode): ) else: ocs_extra = extra_params.pop("ocs", {}) + ocs_verbose = bool(ocs_extra.pop("verbose", False)) override_extra = ocs_extra.pop("override", {}) - immiscible_extra = override_extra.pop("immiscible", {}) + immiscible_extra = ocs_extra.pop("immiscible", {}) overridecfg_kwargs |= override_extra + filterdefs = immiscible_extra.pop("filters", {}) + immisciblecfg_kwargs |= { + f"filter_{k}": filtering.make_filter(filterdefs[k]) + for k in ("reference", "noise", "result", "postcfg") + if k in filterdefs + } + if "filter_reference" in immisciblecfg_kwargs: + immisciblecfg_kwargs["reference"] = "expression" immisciblecfg_kwargs |= immiscible_extra sampler_kwargs |= extra_params + if latent_ref is not None: + overridecfg_kwargs["latent_ref"] = latent_ref["samples"].to( + dtype=torch.float32, device="cpu", copy=True + ) if immiscible_reference == "latent": - if latent_reference is None: + if latent_ref is None: raise ValueError( "latent_ref input must be connected when reference mode is latent" ) - overridecfg_kwargs["latent_ref"] = latent_reference["samples"].to( - dtype=torch.float32, device="cpu", copy=True - ) elif immiscible_reference == "noise": if custom_noise_ref is None: raise ValueError( @@ -1199,6 +1315,7 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode): sampler_function = functools.partial( self.sampler_function, ocs_override_sampler_cfg=OverrideSamplerConfig( + verbose=ocs_verbose, sampler=sampler, sampler_kwargs=sampler_kwargs, immiscible=ImmiscibleConfig(**immisciblecfg_kwargs), @@ -1224,246 +1341,377 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode): @torch.no_grad() def sampler_function( model, - x, - sigmas, + x: torch.Tensor, + sigmas: torch.Tensor, *args: list, ocs_override_sampler_cfg: dict[str] | None = None, noise_sampler=None, extra_args: dict[str] | None = None, **kwargs: dict[str], - ): + ) -> torch.Tensor: cfg = ocs_override_sampler_cfg if cfg is None: raise ValueError("Override sampler config missing!") icfg = cfg.immiscible + if cfg.verbose: + tqdm.write(f"* OCS: Using immiscible config: {icfg}") if extra_args is None: extra_args = {} sig = inspect.signature(cfg.sampler.sampler_function) params = frozenset(cfg.force_params) | sig.parameters.keys() kwargs |= {k: v for k, v in cfg.sampler_kwargs.items() if k in params} - if "noise_sampler" in params: - seed = extra_args.get("seed") - seed_gen = torch.Generator(device="cpu" if cfg.cpu_noise else x.device) - seed_gen.manual_seed(seed if seed is not None else 0) - model_sampling = model.inner_model.inner_model.model_sampling - orig_noise_sampler = kwargs.pop( - "noise_sampler", lambda *_args, **_kwargs: torch.randn_like(x) + if "noise_sampler" not in params and not icfg.force: + return cfg.sampler.sampler_function( + model, + x, + sigmas, + *args, + extra_args=extra_args, + **kwargs, ) - if ( - cfg.custom_noise is not None - and cfg.noise_start_time < 1 - and cfg.noise_end_time > 0 - ): - sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() - custom_noise_sampler = cfg.custom_noise.make_noise_sampler( - x, - sigma_min, - sigma_max, - seed=seed, - cpu=cfg.cpu_noise, - normalized=False, - ) + requires_uncond_modes = { + "uncond", + "cond_sub_uncond", + "denoised_sub_uncond", + } + requires_prev_modes = { + "model_input_prev", + "model_input_sub_model_input_prev", + "noise_gen_prev", + "noise_sub_noise_prev", + "cond_prev", + "uncond_prev", + "denoised_prev", + "denoised_sub_denoised_prev", + } + args_prev: dict | None = None + noise_prev = None + x_orig = x.clone() + seed = extra_args.get("seed") + seed_gen = torch.Generator(device="cpu" if cfg.cpu_noise else x.device) + seed_gen.manual_seed(seed if seed is not None else 0) + if icfg.use_triton and icfg.shuffle_seed is not None: + shuffle_gen = torch.Generator(device="cpu" if cfg.cpu_noise else x.device) + shuffle_gen.manual_seed(icfg.shuffle_seed) + else: + shuffle_gen = None + model_sampling = model.inner_model.inner_model.model_sampling + orig_noise_sampler = kwargs.pop( + "noise_sampler", lambda *_args, **_kwargs: torch.randn_like(x) + ) + if ( + cfg.custom_noise is not None + and cfg.noise_start_time < 1 + and cfg.noise_end_time > 0 + ): + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() + custom_noise_sampler = cfg.custom_noise.make_noise_sampler( + x, + sigma_min, + sigma_max, + seed=seed, + cpu=cfg.cpu_noise, + normalized=False, + ) + else: + custom_noise_sampler = None + sigma_start = model_sampling.percent_to_sigma(cfg.noise_start_time) + sigma_end = model_sampling.percent_to_sigma(cfg.noise_end_time) + + def override_noise_sampler(s, sn, *args, **kwargs): + nonlocal noise_prev + if custom_noise_sampler is None or not sigma_end <= s.max() <= sigma_start: + noise = orig_noise_sampler(s, sn, *args, **kwargs) else: - custom_noise_sampler = None - if custom_noise_sampler is not None: - sigma_start = model_sampling.percent_to_sigma(cfg.noise_start_time) - sigma_end = model_sampling.percent_to_sigma(cfg.noise_end_time) + noise = custom_noise_sampler(s, sn, *args, **kwargs) + noise_prev = noise.clone() + return noise - def override_noise_sampler(s, sn, *args, **kwargs): - if not sigma_end <= s.max() <= sigma_start: - return orig_noise_sampler(s, sn, *args, **kwargs) - return custom_noise_sampler(s, sn, *args, **kwargs) + def fallback_noise_sampler( + s: torch.Tensor, sn: torch.Tensor, *args: list, **kwargs: dict + ) -> torch.Tensor: + nonlocal noise_prev + noise = scale_noise( + override_noise_sampler(s, sn, *args, **kwargs), + normalized=cfg.normalize, + ) + noise_prev = noise.clone() + return noise + if icfg.force or ( + icfg.start_time < 1 + and icfg.end_time > 0 + and (icfg.size > 1 if icfg.batching == "batch" else icfg.size > 0) + and (icfg.blend_mode != "lerp" or icfg.blend != 0) + ): + isigma_start = model_sampling.percent_to_sigma(icfg.start_time) + isigma_end = model_sampling.percent_to_sigma(icfg.end_time) + if icfg.batching.startswith("cycle_"): + ibatching = icfg.batching.split("_")[1:] else: - override_noise_sampler = orig_noise_sampler - if ( - icfg.start_time < 1 - and icfg.end_time > 0 - and (icfg.size > 1 if icfg.batching == "batch" else icfg.size > 0) - and icfg.blend != 0 - ): - isigma_start = model_sampling.percent_to_sigma(icfg.start_time) - isigma_end = model_sampling.percent_to_sigma(icfg.end_time) - if icfg.batching.startswith("cycle_"): - ibatching = icfg.batching.split("_")[1:] - else: - ibatching = icfg.batching - immiscible = ImmiscibleNoise( - size=icfg.size, - batching=ibatching if isinstance(ibatching, str) else "channel", - maximize=icfg.maximize, - distance_scale=icfg.distance_scale, - distance_scale_ref=icfg.distance_scale_ref, + ibatching = icfg.batching + immiscible = ImmiscibleNoise( + size=icfg.size, + batching=ibatching if isinstance(ibatching, str) else "channel", + maximize=icfg.maximize, + distance_scale=icfg.distance_scale, + distance_scale_ref=icfg.distance_scale_ref, + abs_mode=icfg.abs_mode, + abs_distance_mode=icfg.abs_distance_mode, + use_triton=icfg.use_triton, + generator=shuffle_gen, + ) + blend_function = filtering.BLENDING_MODES[icfg.blend_mode] + + ref_latent = None + + requires_patch = icfg.filter_reference or icfg.reference not in { + "latent", + "noise", + "noise_gen_prev", + "initial_latent", + } + + filter_refs = None + using_filters = ( + icfg.filter_reference + or icfg.filter_noise + or icfg.filter_result + or icfg.filter_postcfg + ) + + immiscible_counter = 0 + + if requires_patch: + requires_uncond = ( + using_filters + or icfg.filter_noise + or icfg.filter_result + or icfg.reference in requires_uncond_modes ) - blend_function = filtering.BLENDING_MODES[icfg.blend_mode] - ref_latent = None + def ng_prev_handler(*args, **kwargs): + nonlocal noise_prev + return noise_prev - requires_patch = icfg.reference not in {"latent", "noise"} + ref_handlers = { + "initial_latent": x_orig, + "noise_gen_prev": ng_prev_handler, + "cond": "cond_denoised", + "uncond": "uncond_denoised", + "denoised": "denoised", + "model_input": "input", + "noise_prediction": lambda args: args["input"] - args["denoised"], + "cond_sub_uncond": lambda args: ( + args["cond_denoised"] - args["uncond_denoised"] + ), + "uncond_sub_cond": lambda args: ( + args["uncond_denoised"] - args["cond_denoised"] + ), + "denoised_sub_uncond": lambda args: ( + args["denoised"] - args["uncond_denoised"] + ), + } + if ( + icfg.filter_reference is None + and (ref_handler := ref_handlers.get(icfg.reference)) is None + ): + raise ValueError("Bad immiscible reference type") - if requires_patch: - requires_uncond = icfg.reference in {"uncond", "cond_sub_uncond"} - ref_handlers = { - "cond": "cond_denoised", - "uncond": "uncond_denoised", - "denoised": "denoised", - "model_input": "input", - "noise_prediction": lambda args: args["input"] - - args["denoised"], - "cond_sub_uncond": lambda args: args["cond_denoised"] - - args["uncond_denoised"], - "uncond_sub_cond": lambda args: args["uncond_denoised"] - - args["cond_denoised"], - } - if (ref_handler := ref_handlers.get(icfg.reference)) is None: - raise ValueError("Bad immiscible reference type") + def postcfg(args: dict) -> torch.Tensor: + nonlocal ref_latent, filter_refs, immiscible_counter - def postcfg(args): - nonlocal ref_latent + denoised = args["denoised"] + if using_filters: + uncond = args.get("uncond_denoised") + filter_refs = filtering.FilterRefs( + kvs={ + "immiscible_counter": immiscible_counter, + "sigmas": sigmas.clone(), + "model_sigma": args["sigma"].clone(), + "model_sigma_float": args["sigma"].max().item(), + "x": args["input"].clone(), + "denoised": denoised.clone(), + "cond": args["cond_denoised"].clone(), + "uncond": uncond.clone() + if uncond is not None + else None, + "latent_ref": None + if cfg.latent_ref is None + else cfg.latent_ref.to( + device=denoised.device, + dtype=denoised.dtype, + copy=True, + ), + } + ) + if icfg.filter_reference: + ref_latent = icfg.filter_reference.apply( + denoised, refs=filter_refs + ) + else: ref_latent = ( args.get(ref_handler) if isinstance(ref_handler, str) else ref_handler(args) ) - return args["denoised"] + if icfg.filter_postcfg: + return icfg.filter_postcfg.apply(denoised, refs=filter_refs) + return denoised - extra_args = extra_args | { - "model_options": set_model_options_post_cfg_function( - extra_args.get("model_options", {}).copy(), - postcfg, - disable_cfg1_optimization=requires_uncond, - ) - } + extra_args = extra_args | { + "model_options": set_model_options_post_cfg_function( + extra_args.get("model_options", {}).copy(), + postcfg, + disable_cfg1_optimization=requires_uncond, + ) + } - if icfg.reference == "noise": - ns_ref = cfg.custom_noise_ref.make_noise_sampler( - x, - sigma_min, - sigma_max, - seed=torch.randint( - 0, - 1 << 32, - (1,), - device="cpu" if cfg.cpu_noise else x.device, - dtype=torch.int64, - generator=seed_gen, - ) - .detach() - .cpu() - .item(), - cpu=cfg.cpu_noise, - normalized=False, + if icfg.reference == "noise": + ns_ref = cfg.custom_noise_ref.make_noise_sampler( + x, + sigma_min, + sigma_max, + seed=torch.randint( + 0, + 1 << 32, + (1,), + device="cpu" if cfg.cpu_noise else x.device, + dtype=torch.int64, + generator=seed_gen, ) - elif icfg.reference == "latent": - ref_latent = cfg.latent_ref.to( - dtype=x.dtype, device=x.device, copy=True + .detach() + .cpu() + .item(), + cpu=cfg.cpu_noise, + normalized=False, + ) + elif icfg.reference == "latent": + ref_latent = cfg.latent_ref.to( + dtype=x.dtype, device=x.device, copy=True + ) + if icfg.norm_ref_scale != 0: + ref_latent = scale_noise( + ref_latent, icfg.norm_ref_scale, normalized=True ) - if icfg.norm_ref_scale != 0: - ref_latent = scale_noise( - ref_latent, icfg.norm_ref_scale, normalized=True - ) - if ref_latent.shape[1:] != x.shape[1:]: + if ref_latent.shape[1:] != x.shape[1:]: + raise ValueError( + "Reference latent shape must match shape of generation with exception that batch size may be 1", + ) + if ref_latent.shape[0] != x.shape[0]: + if ref_latent.shape[0] == 1: + ref_latent = ref_latent.expand(x.shape) + else: raise ValueError( - "Reference latent shape must match shape of generation with exception that batch size may be 1", + "Reference latent batch size must be either 1 or equal to generation batch size" ) - if ref_latent.shape[0] != x.shape[0]: - if ref_latent.shape[0] == 1: - ref_latent = ref_latent.expand(x.shape) - else: - raise ValueError( - "Reference latent batch size must be either 1 or equal to generation batch size" - ) - if icfg.blend != 1 and cfg.custom_noise_blend is not None: - ns_blend = cfg.custom_noise_blend.make_noise_sampler( - x, - sigma_min, - sigma_max, - seed=torch.randint( - 0, - 1 << 32, - (1,), - device="cpu" if cfg.cpu_noise else x.device, - dtype=torch.int64, - generator=seed_gen, - ) - .detach() - .cpu() - .item(), - cpu=cfg.cpu_noise, - normalized=False, + if icfg.blend != 1 and cfg.custom_noise_blend is not None: + ns_blend = cfg.custom_noise_blend.make_noise_sampler( + x, + sigma_min, + sigma_max, + seed=torch.randint( + 0, + 1 << 32, + (1,), + device="cpu" if cfg.cpu_noise else x.device, + dtype=torch.int64, + generator=seed_gen, ) - else: - ns_blend = None - - immiscible_counter = 0 - - def noise_sampler(s, sn, *args, **kwargs): - nonlocal ref_latent, immiscible_counter - if not isigma_end <= s.max() <= isigma_start: - return override_noise_sampler(s, sn, *args, **kwargs) - if icfg.reference == "noise": - ref_latent = ns_ref(s, sn) - elif ref_latent is None: - raise ValueError("Immiscible reference type not available") - if not isinstance(ibatching, str): - immiscible.batching = ibatching[ - immiscible_counter % len(ibatching) - ] - immiscible_counter += 1 - blend_in_batch = icfg.blend != 1 and ns_blend is None - # blending = icfg.blend != 1 - if icfg.norm_ref_scale != 0 and icfg.reference != "latent": - ref_latent = scale_noise( - ref_latent, icfg.norm_ref_scale, normalized=True - ) - batch_size = ref_latent.shape[0] - noise_batch = torch.cat( - tuple( - override_noise_sampler(s, sn) - for _ in range(max(1, icfg.size) + int(blend_in_batch)) - ) - ) - immiscible_noise = immiscible.unbatch( - immiscible.immiscible( - immiscible.batch( - scale_noise( - noise_batch[batch_size * int(blend_in_batch) :], - 1.0 - if icfg.norm_noise_scale == 0 - else icfg.norm_noise_scale, - normalized=icfg.norm_noise_scale != 0, - ) - ), - immiscible.batch(ref_latent), - ), - ref_latent.shape, - ) - immiscible_noise = scale_noise( - immiscible_noise, normalized=cfg.normalize - ) - if icfg.blend != 1: - immiscible_noise = blend_function( - noise_batch[:batch_size] - if blend_in_batch - else ns_blend(s, sn), - immiscible_noise, - icfg.blend, - ) - return scale_noise(immiscible_noise, normalized=cfg.normalize) - + .detach() + .cpu() + .item(), + cpu=cfg.cpu_noise, + normalized=False, + ) else: - if not cfg.normalize: - noise_sampler = override_noise_sampler - else: + ns_blend = None - def noise_sampler(s, sn, *args, **kwargs): - return scale_noise( - override_noise_sampler(s, sn, *args, **kwargs), - normalized=True, - ) + def noise_sampler( + s: torch.Tensor, + sn: torch.Tensor, + *args: list, + **kwargs: dict, + ) -> torch.Tensor: + nonlocal ref_latent, filter_refs, immiscible_counter, noise_prev + if not isigma_end <= s.max() <= isigma_start: + return override_noise_sampler(s, sn, *args, **kwargs) + if icfg.filter_noise or icfg.filter_result: + curr_refs = filtering.FilterRefs( + kvs={ + "sigma": s.clone(), + "sigma_next": sn.clone(), + "immiscible_counter": immiscible_counter, + } + ) + if filter_refs is not None: + curr_refs = curr_refs | filter_refs + if icfg.reference == "noise": + # FIXME: This doesn't honor immiscible ref scale + ref_latent = ns_ref(s, sn) + elif icfg.reference == "initial_latent": + ref_latent = x_orig + elif icfg.reference == "noise_gen_prev": + if noise_prev is None: + return fallback_noise_sampler(s, sn, *args, **kwargs) + ref_latent = noise_prev + elif ref_latent is None: + raise ValueError("Immiscible reference type not available") + if not isinstance(ibatching, str): + immiscible.batching = ibatching[immiscible_counter % len(ibatching)] + immiscible_counter += 1 + blend_in_batch = icfg.blend != 1 and ns_blend is None + # blending = icfg.blend != 1 + if icfg.norm_ref_scale != 0 and icfg.reference != "latent": + ref_latent = scale_noise( + ref_latent, icfg.norm_ref_scale, normalized=True + ) + batch_size = ref_latent.shape[0] + noise_batch = torch.cat( + tuple( + override_noise_sampler(s, sn) + for _ in range(max(1, icfg.size) + int(blend_in_batch)) + ) + ) + if icfg.filter_noise: + noise_batch = icfg.filter_noise.apply(noise_batch, refs=curr_refs) + immiscible_noise = immiscible.unbatch( + immiscible.immiscible( + immiscible.batch( + scale_noise( + noise_batch[batch_size * int(blend_in_batch) :], + 1.0 + if icfg.norm_noise_scale == 0 + else icfg.norm_noise_scale, + normalized=icfg.norm_noise_scale != 0, + ) + ), + immiscible.batch(ref_latent), + ), + ref_latent.shape, + ) + immiscible_noise = scale_noise( + immiscible_noise, normalized=cfg.normalize + ) + if icfg.blend != 1: + immiscible_noise = blend_function( + noise_batch[:batch_size] if blend_in_batch else ns_blend(s, sn), + immiscible_noise, + icfg.blend, + ) + if icfg.filter_result: + immiscible_noise = icfg.filter_result.apply( + immiscible_noise, refs=curr_refs + ) + noise = scale_noise(immiscible_noise, normalized=cfg.normalize) + noise_prev = noise.clone() + return noise - kwargs["noise_sampler"] = noise_sampler + else: + noise_sampler = fallback_noise_sampler + + kwargs["noise_sampler"] = noise_sampler return cfg.sampler.sampler_function( model, x, @@ -1481,34 +1729,75 @@ class ExpressionFilteredNoiseItem(CustomNoiseItemBase): *, normalize, noise: object, - noise_filter, + noise_filter: filtering.Filter, + ref_filter: filtering.Filter | None, + latent_refs: dict, ): super().__init__( factor, noise=noise, noise_filter=noise_filter, + ref_filter=ref_filter, normalize=normalize, + latent_refs={k: v.clone() for k, v in latent_refs.items()}, ) def clone_key(self, k): if k == "noise": return self.noise.clone() + if k == "latent_refs": + return {k: v.clone() for k, v in self.latent_refs.items()} return super().clone_key(k) - def make_noise_sampler(self, x, *args, normalized=True, **kwargs): + def make_noise_sampler( + self, + x: torch.Tensor, + *args: Any, + normalized=True, + **kwargs: Any, + ) -> Callable: noise_filter = self.noise_filter + initial_shape = x.shape + latent_refs = self.latent_refs + if self.ref_filter is not None: + x = self.ref_filter.apply( + x, + refs=filtering.FilterRefs({k: v.to(x) for k, v in latent_refs.items()}), + ) ns = self.noise.make_noise_sampler(x, *args, normalized=False, **kwargs) normalize_noise = self.normalize != False and normalized # noqa: E712 factor = self.factor + sample_counter = 0 + initial_x = x.clone() + last_noise = None - def noise_sampler(s, sn, *args, **kwargs): + def noise_sampler(s, sn, *args: Any, **kwargs: Any) -> torch.Tensor: + nonlocal sample_counter, last_noise noise = ns(s, sn) - refs = filtering.FilterRefs({ - "sigma": s.clone() if isinstance(s, torch.Tensor) else s, - "sigma_next": sn.clone() if isinstance(sn, torch.Tensor) else sn, - }) + if ( + isinstance(s, torch.Tensor) + and isinstance(sn, torch.Tensor) + and s.ndim < initial_x.ndim + ): + padded_shape = tuple(-1 if d == 0 else 1 for d in range(initial_x.ndim)) + s, sn = s.reshape(padded_shape), sn.reshape(padded_shape) + + refs = filtering.FilterRefs( + { + "sigma": s.clone() if isinstance(s, torch.Tensor) else s, + "sigma_next": sn.clone() if isinstance(sn, torch.Tensor) else sn, + "sample_counter": sample_counter, + "initial_x": initial_x, + "initial_shape": initial_shape, + "last_noise": last_noise, + } + | {k: v.to(noise) for k, v in latent_refs.items()} + ) + sample_counter += 1 noise = noise_filter.apply(noise, refs=refs) - return scale_noise(noise, factor, normalized=normalize_noise) + noise = scale_noise(noise, factor, normalized=normalize_noise) + last_noise = noise.clone() + return noise return noise_sampler @@ -1546,6 +1835,13 @@ class ExpressionFilteredNoiseNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): }, ), } + if "optional" not in result: + result["optional"] = {} + result["optional"] |= { + "latent_ref_1_opt": ("LATENT",), + "latent_ref_2_opt": ("LATENT",), + "latent_ref_3_opt": ("LATENT",), + } return result @classmethod @@ -1560,20 +1856,40 @@ class ExpressionFilteredNoiseNode(CustomNoiseNodeBase, NormalizeNoiseNodeMixin): normalize: str, custom_noise: object, yaml_config: str, + latent_ref_1_opt: dict | None = None, + latent_ref_2_opt: dict | None = None, + latent_ref_3_opt: dict | None = None, ) -> tuple: config = yaml.safe_load(yaml_config) if not isinstance(config, dict) or "filter" not in config: raise ValueError( "Bad YAML config type (must be object) or missing filter key in config" ) - filter_def = config.get("filter") - if not isinstance(filter_def, dict): - raise ValueError("Bad type for filter definition, must be object") - ocs_filter = filtering.make_filter(filter_def) + noise_filter_def = config.get("filter") + if not isinstance(noise_filter_def, dict): + raise TypeError("Bad type for filter definition, must be object") + ref_filter_def = config.get("ref_filter") + if ref_filter_def is not None and not isinstance(ref_filter_def, dict): + raise TypeError("ref_filter key must be an object if present") + latent_refs = { + k: v["samples"].to(device="cpu", dtype=torch.float32, copy=True) + for k, v in ( + ("latent_ref_1", latent_ref_1_opt), + ("latent_ref_2", latent_ref_2_opt), + ("latent_ref_3", latent_ref_3_opt), + ) + if v is not None + } + noise_filter = filtering.make_filter(noise_filter_def) + ref_filter = ( + None if ref_filter_def is None else filtering.make_filter(ref_filter_def) + ) return super().go( factor, rescale=rescale, normalize=self.get_normalize(normalize), noise=custom_noise.clone(), - noise_filter=ocs_filter, + noise_filter=noise_filter, + ref_filter=ref_filter, + latent_refs=latent_refs, ) diff --git a/py/custom_noise/noise_perlin.py b/py/custom_noise/noise_perlin.py index 8ff023e..c70df35 100644 --- a/py/custom_noise/noise_perlin.py +++ b/py/custom_noise/noise_perlin.py @@ -1,412 +1,149 @@ -import itertools +# Initial revision based on Perlin generation routines from https://github.com/Extraltodeus/noise_latent_perlinpinpin which was based on https://gist.github.com/vadimkantorov/ac1b097753f217c5c11bc2ff396e0a57 which was based on https://github.com/pvigier/perlin-numpy import math +from typing import Any, Callable, NamedTuple, Sequence import torch - from comfy import model_management +from tqdm import tqdm from .. import filtering from ..latent import normalize_to_scale from ..noise import scale_noise from .base import CustomNoiseItemBase, NormalizeNoiseNodeMixin -# Perlin generation routines based on https://github.com/Extraltodeus/noise_latent_perlinpinpin which was based on https://gist.github.com/vadimkantorov/ac1b097753f217c5c11bc2ff396e0a57 which was based on https://github.com/pvigier/perlin-numpy - def smoothstep_function(t): return 6 * t**5 - 15 * t**4 + 10 * t**3 -class DEFAULTS: - depth = 16 - res = ((1,), (1,), (1,)) - octaves = 2 - persistence = (1.0,) - lacunarity = ((2,), (2,), (2,)) - initial_amplitude = 1.0 - initial_frequency = (1.0, 1.0, 1.0) - break_pattern = 1.0 - detail_level = 0.0 - tileable = (False, False, False) - fade = smoothstep_function - blend = "lerp" - pattern_break_blend = "lerp" - depth_over_channels = False - initial_depth = 0 - wrap_depth = 0 - max_depth = -1 - pad = (0, 0, 0) - generator = None - device = "default" +class BlendFunction(NamedTuple): + name: str = "lerp" + blend_function: Callable[..., torch.Tensor] = torch.lerp + + def __call__(self, *args: Any, **kwargs: Any) -> torch.Tensor: + return self.blend_function(*args, **kwargs) + + +class Perlin(NamedTuple): + depth: int = 16 + res: tuple[tuple[float, ...], ...] = ((1.0,), (1.0,), (1.0,)) + octaves: int = 2 + persistence: tuple[float, ...] = (1.0,) + lacunarity: tuple[tuple[float, ...], ...] = ((2,), (2,), (2,)) + initial_amplitude: float = 1.0 + initial_frequency: tuple[float, ...] = (1.0, 1.0, 1.0) + break_pattern: float = 0.99 + break_pattern_multiplier: float = 100000.0 + break_pattern_use_frac: bool = True + detail_level: float = 0.0 + ridge_blend: BlendFunction = BlendFunction() + ridge_weight: float = 0.0 + ridge_scale: float = 1.0 + warp_strength: float = 0.0 + octave_shift: float = 0.0 + curl_strength: float = 0.0 + curl_dims: tuple[int, int] = (0, 1) + tileable: tuple[bool, ...] = (False, False, False) + fade: Callable[[torch.Tensor], torch.Tensor] = smoothstep_function + blend: BlendFunction = BlendFunction() + pattern_break_blend: BlendFunction = BlendFunction() + depth_over_channels: bool = False + initial_depth: int = 0 + wrap_depth: int = 0 + max_depth: int = -1 + pad: tuple[int, ...] = (0, 0, 0) + pad_mode: str = "replicate" + generator: torch.Generator | None = None + device: str | torch.device = "default" + dtype: torch.dtype | None = None @classmethod - def get_commasep(cls, key, idx=None): - val = getattr(cls, key) - if idx is not None: - val = val[idx] - return ", ".join(repr(v) for v in val) - - -def rand_perlin( - shape, - res, - *, - tileable=DEFAULTS.tileable, - fade=DEFAULTS.fade, - blend=torch.lerp, - generator=DEFAULTS.generator, - device=DEFAULTS.device, -): - dims = len(res) - didxs = tuple(range(dims)) - delta, d = zip(*((res[i] / shape[i], int(round(shape[i] / res[i]))) for i in didxs)) - - grid = ( - torch.stack( - torch.meshgrid(*(torch.arange(0, res[i], delta[i]) for i in didxs)), - dim=-1, - ) - % 1 - ).to(device=device) - - noise = ( - 2 - * math.pi - * torch.rand( - max(1, dims - 1), - *(round(res[i]) + 1 for i in didxs), - generator=generator, - device=device, - ) - ) - if dims == 1: - gradients = torch.cos(noise[0]) - elif dims == 2: - gradients = torch.stack((torch.cos(noise[0]), torch.sin(noise[0])), dim=-1) - elif dims == 3: - gradients = torch.stack( - ( - torch.sin(noise[0]) * torch.cos(noise[1]), - torch.sin(noise[0]) * torch.sin(noise[1]), - torch.cos(noise[0]), - ), - dim=-1, - ) - elif dims == 4: - # No idea if this makes sense. - gradients = torch.stack( - ( - torch.sin(noise[0]) * torch.cos(noise[1]), - torch.sin(noise[0]) * torch.sin(noise[1]), - torch.sin(noise[1]) * torch.cos(noise[2]), - torch.sin(noise[1]) * torch.sin(noise[2]), - ), - dim=-1, - ) - else: - raise ValueError("Currently only dimensions up to 4 are supported") - del noise - - for tidx, tile in enumerate(tileable[:dims]): - if not tile: - continue - gradients[tuple(-1 if didx == tidx else None for didx in didxs)] = gradients[ - tuple(0 if didx == tidx else None for didx in didxs) - ] - - shape_slices = tuple(slice(0, shape[i]) for i in didxs) - - def tile_grads(slices): - result = gradients[tuple(slice(*slices[i]) for i in didxs)] - for i in didxs: - result = result.repeat_interleave(d[i], i) - return result - - def dot(grad, shift): - return ( - torch.stack( - tuple(grid[(*shape_slices, i)] + shift[i] for i in didxs), dim=-1 + def build(cls, **kwargs: Any) -> "Perlin": + dfl = cls() + depth = kwargs.get("depth", dfl.depth) + for bk in ("blend", "pattern_break_blend", "ridge_blend"): + bv = kwargs.pop(bk, getattr(dfl, bk)) + kwargs[bk] = ( + BlendFunction(bv, filtering.BLENDING_MODES[bv]) + if isinstance(bv, str) + else bv ) - * grad[shape_slices] - ).sum(dim=-1) - - # It's just binary with the bits reversed and -1 for enabled columns. - def get_shift(n, dims, *, on_value, off_value): - return tuple( - on_value if n & (1 << bitidx) else off_value for bitidx in range(dims) - ) - - def blend_reduce(vals, t, depth=0): - curr_t = t[..., depth] - if len(vals) == 2: - return blend(*vals, curr_t) - return blend_reduce( - tuple(blend(v1, v2, curr_t) for v1, v2 in itertools.batched(vals, 2)), - t, - depth + 1, - ) - - ns = tuple( - dot( - tile_grads(get_shift(i, dims, off_value=(None, -1), on_value=(1, None))), - get_shift(i, dims, off_value=0, on_value=-1), - ) - for i in range(1 << dims) - ) - return math.sqrt(2) * blend_reduce(ns, fade(grid[shape_slices])) - - -def generate_fractal_noise( - shape, - res=DEFAULTS.res, - octaves=DEFAULTS.octaves, - persistence=DEFAULTS.persistence, - lacunarity=DEFAULTS.lacunarity, - initial_amplitude=DEFAULTS.initial_amplitude, - initial_frequency=DEFAULTS.initial_frequency, - tileable=DEFAULTS.tileable, - fade=DEFAULTS.fade, - blend=torch.lerp, - generator=DEFAULTS.generator, - device=DEFAULTS.device, -): - ndim = len(shape) - - def get_wrap_dim(val, *dims): - for dim in dims: - nelem = len(val) if not isinstance(val, torch.Tensor) else val.shape[0] - val = val[dim % nelem] - return val - - def get_unwrapped_octaves_dims(val): - return torch.tensor( - tuple( - get_wrap_dim(val, didx, oidx) - for oidx in range(octaves) - for didx in range(ndim) - ), - dtype=torch.float, - device="cpu", - ).view(octaves, ndim) - - res = get_unwrapped_octaves_dims(res) - lacunarity = get_unwrapped_octaves_dims(lacunarity) - initial_frequency = initial_frequency[-ndim:] - persistence = persistence[:octaves] - noise = torch.zeros(shape, dtype=torch.float32, device=device) - frequency = torch.ones(ndim, dtype=torch.float, device="cpu") - frequency[: len(initial_frequency)] = frequency.new(initial_frequency) - amplitude = initial_amplitude - - for octave in range(octaves): - noise += amplitude * rand_perlin( - shape, - tuple( - frequency[didx].item() * res[octave][didx].item() - for didx in range(ndim) - ), - tileable=tileable, - fade=fade, - blend=blend, - generator=generator, - device=device, - ) - # print( - # f"Octave {octave}: freq={frequency}, amp={amplitude}, lac={lacunarity[octave]}, pers={get_wrap_dim(persistence, octave)}" - # ) - frequency *= lacunarity[octave] - amplitude *= get_wrap_dim(persistence, octave) - # print(f"Octave {octave}: POST: freq={frequency}, amp={amplitude}") - return noise - - -def create_noisy_latents_perlin( - width, - height, - depth, - *, - batch_size=1, - detail_level=DEFAULTS.detail_level, - octaves=DEFAULTS.octaves, - persistence=DEFAULTS.persistence, - lacunarity=DEFAULTS.lacunarity, - tileable=DEFAULTS.tileable, - res=DEFAULTS.res, - break_pattern=DEFAULTS.break_pattern, - channels=4, - blend=torch.lerp, - pattern_break_blend=torch.lerp, - depth_over_channels=DEFAULTS.depth_over_channels, - pad=DEFAULTS.pad, - initial_frequency=DEFAULTS.initial_frequency, - initial_amplitude=DEFAULTS.initial_amplitude, - generator=DEFAULTS.generator, - device=DEFAULTS.device, -): - pad_depth, pad_height, pad_width = pad - if depth < 1: - depth_over_channels = False - pad_depth = 0 - shape = (height, width) - eff_shape = ( - height + pad_height * 2, - width + pad_width * 2, - ) - eff_channels = channels if not depth_over_channels else 1 - eff_depth = depth if not depth_over_channels else depth * channels - if depth > 0: - shape = (depth, height, width) - eff_shape = ( - eff_depth + pad_depth * 2, - height + pad_height * 2, - width + pad_width * 2, - ) - noise = torch.zeros( - (batch_size, channels, *shape), - dtype=torch.float32, - device=device, - ) - noise_dims = len(shape) - for i in range(batch_size): - for j in range(eff_channels): - noise_values = generate_fractal_noise( - eff_shape, - res=res, - octaves=octaves, - persistence=persistence, - lacunarity=lacunarity, - tileable=tileable, - blend=blend, - initial_frequency=initial_frequency, - initial_amplitude=initial_amplitude, - generator=generator, - device=device, - ) - noise_values = normalize_to_scale(noise_values, -1.0, 1.0, dim=()) - if break_pattern != 0: - result = torch.remainder(torch.abs(noise_values) * 1000000, 11) / 11 - result = ( - ((1 + detail_level / 10) * torch.erfinv(2 * result - 1) * (2**0.5)) - .mul_(0.2) - .clamp_(-1, 1) - ) - result = pattern_break_blend(noise_values, result, break_pattern) - else: - result = noise_values - if pad_width + pad_height + pad_depth > 0: - result = ( - result[ - ..., - pad_depth : eff_depth + pad_depth, - pad_height : height + pad_height, - pad_width : width + pad_width, - ] - if noise_dims == 3 - else result[ - ..., - pad_height : height + pad_height, - pad_width : width + pad_width, - ] - ) - if not depth_over_channels: - noise[i, j, ...] = result - continue - noise[i, ...] = result.view(depth, channels, height, width).movedim(0, 1) - return noise.movedim(-3, 0) if noise_dims == 3 else noise - - -class PerlinItem(CustomNoiseItemBase): - def __init__( - self, - factor, - *, - depth=20, - detail_level=DEFAULTS.detail_level, - octaves=DEFAULTS.octaves, - persistence=DEFAULTS.persistence, - lacunarity_depth=DEFAULTS.lacunarity[0], - lacunarity_height=DEFAULTS.lacunarity[1], - lacunarity_width=DEFAULTS.lacunarity[2], - lacunarity=None, - tileable_depth=DEFAULTS.tileable[0], - tileable_height=DEFAULTS.tileable[1], - tileable_width=DEFAULTS.tileable[2], - tileable=None, - res_depth=DEFAULTS.res[0], - res_height=DEFAULTS.res[1], - res_width=DEFAULTS.res[2], - res=None, - initial_frequency_depth=DEFAULTS.initial_frequency[0], - initial_frequency_height=DEFAULTS.initial_frequency[1], - initial_frequency_width=DEFAULTS.initial_frequency[2], - initial_frequency=None, - initial_amplitude=DEFAULTS.initial_amplitude, - wrap_depth=DEFAULTS.wrap_depth, - initial_depth=DEFAULTS.initial_depth, - max_depth=DEFAULTS.max_depth, - break_pattern=DEFAULTS.break_pattern, - blend=DEFAULTS.blend, - pattern_break_blend=DEFAULTS.pattern_break_blend, - depth_over_channels=DEFAULTS.depth_over_channels, - pad=None, - pad_depth=DEFAULTS.pad[0], - pad_height=DEFAULTS.pad[1], - pad_width=DEFAULTS.pad[2], - device=None, - normalized=None, - **kwargs, - ): - if tileable is None: - tileable = (tileable_depth, tileable_height, tileable_width)[ - int(depth == 0) : - ] - if res is None: - res = self.maybe_parse_dhw_triple( - (res_depth, res_height, res_width), depth, int - ) - if lacunarity is None: - lacunarity = self.maybe_parse_dhw_triple( + lacunarity = kwargs.pop("lacunarity", None) + kwargs["lacunarity"] = ( + cls.maybe_parse_dhw_triple( ( - lacunarity_depth, - lacunarity_height, - lacunarity_width, + kwargs.pop("lacunarity_depth", dfl.lacunarity[0]), + kwargs.pop("lacunarity_height", dfl.lacunarity[1]), + kwargs.pop("lacunarity_width", dfl.lacunarity[2]), ), depth, ) - if pad is None: - pad = (pad_depth, pad_height, pad_width) - if initial_frequency is None: - initial_frequency = ( - initial_frequency_depth, - initial_frequency_height, - initial_frequency_width, - )[int(depth == 0) :] - persistence = self.maybe_parse_commasep_list(persistence) - super().__init__( - factor, - depth=depth, - detail_level=detail_level, - octaves=octaves, - persistence=persistence, - lacunarity=lacunarity, - tileable=tileable, - res=res, - initial_frequency=initial_frequency, - initial_amplitude=initial_amplitude, - initial_depth=initial_depth, - wrap_depth=wrap_depth, - max_depth=max_depth, - break_pattern=break_pattern, - blend=blend, - pattern_break_blend=pattern_break_blend, - depth_over_channels=depth_over_channels, - pad=pad, - device=device, - normalized=normalized - if not isinstance(normalized, str) - else NormalizeNoiseNodeMixin.get_normalize(normalized), - **kwargs, + if lacunarity is None + else tuple(lacunarity) ) + res = kwargs.pop("res", None) + kwargs["res"] = ( + cls.maybe_parse_dhw_triple( + ( + kwargs.pop("res_depth", dfl.res[0]), + kwargs.pop("res_height", dfl.res[1]), + kwargs.pop("res_width", dfl.res[2]), + ), + depth, + ) + if res is None + else tuple(res) + ) + pad = kwargs.pop("pad", None) + kwargs["pad"] = ( + ( + kwargs.pop("pad_depth", dfl.pad[0]), + kwargs.pop("pad_height", dfl.pad[1]), + kwargs.pop("pad_width", dfl.pad[2]), + ) + if pad is None + else tuple(pad) + ) + initial_frequency = kwargs.pop("initial_frequency", None) + kwargs["initial_frequency"] = ( + ( + kwargs.pop("initial_frequency_depth", dfl.initial_frequency[0]), + kwargs.pop("initial_frequency_height", dfl.initial_frequency[1]), + kwargs.pop("initial_frequency_width", dfl.initial_frequency[2]), + )[int(depth == 0) :] + if initial_frequency is None + else tuple(initial_frequency) + ) + tileable = kwargs.pop("tileable", None) + kwargs["tileable"] = ( + ( + kwargs.pop("tileable_depth", dfl.tileable[0]), + kwargs.pop("tileable_height", dfl.tileable[1]), + kwargs.pop("tileable_width", dfl.tileable[2]), + )[int(depth == 0) :] + if tileable is None + else tuple(tileable) + ) + persistence = kwargs.pop("persistence", None) + if persistence is not None: + kwargs["persistence"] = ( + cls.maybe_parse_commasep_list(persistence) + if isinstance(persistence, str) + else tuple(persistence) + ) + curl_dims = kwargs.pop("curl_dims", None) + if curl_dims is not None: + kwargs["curl_dims"] = tuple( + int(v) + for v in ( + cls.maybe_parse_commasep_list(curl_dims) + if isinstance(curl_dims, str) + else curl_dims + ) + ) + fs = frozenset(cls._fields) + kwargs = {k: v for k, v in kwargs.items() if k in fs} + return cls(**kwargs) @classmethod def maybe_parse_dhw_triple(cls, val, depth, convert=float): @@ -418,6 +155,415 @@ class PerlinItem(CustomNoiseItemBase): return val return tuple(convert(v) for v in val.strip().split(",") if v.strip()) + def get_commasep(self, key, idx=None): + val = getattr(self, key) + if idx is not None: + val = val[idx] + return ", ".join(repr(v) for v in val) + + def octave( + self, + shape: Sequence[int], + res: Sequence[float], + *, + batch_size: int = 1, + channels: int = 1, + octave: int = 0, + base_noise: torch.Tensor | None = None, + warp: torch.Tensor | None = None, + ): + shape = tuple(shape) + res = tuple(res) + dims = len(res) + didxs = tuple(range(dims)) + + coords = tuple( + torch.linspace( + 0, + res[i], + shape[i] + 1, + device=self.device, + dtype=self.dtype, + )[:-1] + for i in didxs + ) + grid_coords = torch.meshgrid(*coords, indexing="ij") + p = torch.stack(grid_coords, dim=-1) + + # Expand `p` to include Batch and Channel dimensions + p = ( + p.unsqueeze(0) + .unsqueeze(0) + .expand( + batch_size, + channels, + *((-1,) * (dims + 1)), + ) + ) + + # Domain Warping: Apply warp before calculating p0 and grid. + if warp is not None and self.warp_strength != 0.0: + p += warp.unsqueeze(-1) * self.warp_strength + + # Now calculate indices and bounds safely + p0 = p.floor().long() + grid = p - p0 + grad_shape = tuple(int(math.ceil(res[i])) + 1 for i in didxs) + + gradients = torch.randn( + batch_size, + channels, + *grad_shape, + dims, + generator=self.generator, + device=self.device, + dtype=self.dtype, + ) + gradients = torch.nn.functional.normalize(gradients, dim=-1) + + # Modulate Perlin amplitude using base noise + if base_noise is not None: + gradients = gradients * base_noise.unsqueeze(-1).to(gradients) + + octave_shift = round((1.0 + octave) * self.octave_shift) + if octave_shift != 0: + gradients = gradients.roll(dims=-1, shifts=octave_shift) + + if dims > 1 and self.curl_strength != 0: + # Generate a random spin angle for every single gradient point on the grid + angles = torch.randn( + batch_size, + channels, + *grad_shape, + generator=self.generator, + device=gradients.device, + dtype=gradients.dtype, + ).mul_(self.curl_strength) + + cos_a = angles.cos() + sin_a = angles.sin_() + + # Grab the first two axes by default (e.g., Depth/Height, or Height/Width) + d1, d2 = self.curl_dims[:2] + g0 = gradients[..., d1].clone() + g1 = gradients[..., d2].clone() + + # Apply 2D Rotation Matrix to twist the vectors + gradients[..., d1] = (g0 * cos_a).sub_(g1 * sin_a) + gradients[..., d2] = (g0 * sin_a).add_(g1 * cos_a) + + def get_shift(n, dims, *, on_value, off_value): + return tuple( + on_value if n & (1 << bitidx) else off_value for bitidx in range(dims) + ) + + def blend_reduce(vals, t, depth=0): + curr_t = t[..., depth] + if len(vals) == 2: + return self.blend(*vals, curr_t) + pairs = zip(vals[0::2], vals[1::2]) + return blend_reduce( + tuple(self.blend(v1, v2, curr_t) for v1, v2 in pairs), t, depth + 1 + ) + + ns = [] + + # Pre-calculate batched indices for tensor indexing + b_idx = torch.arange(batch_size, device=self.device).view( + batch_size, 1, *[1] * dims + ) + c_idx = torch.arange(channels, device=self.device).view( + 1, channels, *[1] * dims + ) + + for i in range(1 << dims): + shift = get_shift(i, dims, off_value=0, on_value=1) + + idx = p0.clone() + for dim in range(dims): + idx[..., dim] += shift[dim] + idx[..., dim] %= grad_shape[dim] - int(self.tileable[dim]) + + spatial_indices = tuple(idx[..., dim] for dim in range(dims)) + grad = gradients[(b_idx, c_idx) + spatial_indices] + + grid_shift = get_shift(i, dims, off_value=0, on_value=-1) + grid_shift_tensor = torch.tensor( + grid_shift, dtype=self.dtype, device=self.device + ) + + d = ((grid + grid_shift_tensor) * grad).sum(dim=-1) + ns.append(d) + + return blend_reduce(ns, self.fade(grid)).mul_(2.0**0.5) + + @staticmethod + def get_wrap_dim(val, *dims): + for dim in dims: + nelem = len(val) if not isinstance(val, torch.Tensor) else val.shape[0] + val = val[dim % nelem] + return val + + def get_unwrapped_octaves_dims(self, val, ndim: int) -> torch.Tensor: + return torch.tensor( + tuple( + self.get_wrap_dim(val, didx, oidx) + for oidx in range(self.octaves) + for didx in range(ndim) + ), + dtype=self.dtype, + device=self.device, + ).reshape(self.octaves, ndim) + + def generate_octaves( + self, + shape: Sequence[int], + *, + batch_size: int = 1, + channels: int = 1, + base_noise: torch.Tensor | None = None, + ) -> torch.Tensor: + shape = tuple(shape) + ndim = len(shape) + + amplitude = self.initial_amplitude + res = self.get_unwrapped_octaves_dims(self.res, ndim) + lacunarity = self.get_unwrapped_octaves_dims(self.lacunarity, ndim) + persistence = self.persistence[: self.octaves] + previous_octave = None + + initial_frequency = self.initial_frequency[-ndim:] + frequency = torch.ones(ndim, dtype=self.dtype, device=self.device) + frequency[: len(initial_frequency)] = frequency.new(initial_frequency) + noise = torch.zeros( + (batch_size, channels, *shape), + dtype=self.dtype, + device=self.device, + ) + + for octave in range(self.octaves): + octave_res = tuple( + frequency[didx].item() * res[octave][didx].item() + for didx in range(ndim) + ) + + grad_shape = tuple(int(math.ceil(octave_res[i])) + 1 for i in range(ndim)) + + if base_noise is None: + octave_base_noise = None + else: + if base_noise.shape[-len(grad_shape) :] != grad_shape: + mode = ( + "bilinear" + if ndim == 2 + else ("trilinear" if ndim == 3 else "nearest") + ) + octave_base_noise = torch.nn.functional.interpolate( + base_noise, + size=grad_shape, + mode=mode, + **( + {"align_corners": False} + if mode not in ("nearest", "area") + else {} + ), + ) + else: + octave_base_noise = base_noise + + octave_output = self.octave( + shape, + octave_res, + batch_size=batch_size, + channels=channels, + octave=octave, + base_noise=octave_base_noise, + warp=previous_octave, + ) + + if self.ridge_weight != 0.0: + ridge = ( + 1.0 + - octave_output.div( + octave_output.abs().max().clamp_min_(1e-07) + ).abs_() + ) + ridge -= 0.5 + ridge *= 2.0 * self.ridge_scale + octave_output = self.ridge_blend( + octave_output, + ridge, + self.ridge_weight, + ) + + noise += amplitude * octave_output + previous_octave = octave_output + + frequency *= lacunarity[octave] + amplitude *= self.get_wrap_dim(persistence, octave) + + return noise + + # Based on approach from https://github.com/Extraltodeus/noise_latent_perlinpinpin + @staticmethod + def break_pattern_func( + t: torch.Tensor, + detail: float = 0.0, + *, + multiplier: float = 1000000.0, + use_frac: bool = False, + clamp_low: float = -5.0, + clamp_high: float = 5.0, + ) -> torch.Tensor: + detail_factor = (1 + detail * 0.1) * 2.0**0.5 * 0.2 + result = t.abs().mul_(multiplier) + result = result.frac_() if use_frac else result.remainer_(11).div_(11) + return ( + result.mul_(2) + .sub_(1) + .erfinv_() + .mul_(detail_factor) + .clamp_(clamp_low, clamp_high) + ) + + def __call__( + self, + width: int, + height: int, + *, + batch_size: int = 1, + channels: int = 4, + base_noise: torch.Tensor | None = None, + ) -> torch.Tensor: + depth = self.depth + pad_depth, pad_height, pad_width = self.pad[:3] + depth_over_channels = self.depth_over_channels + + if depth < 1: + depth_over_channels = False + pad_depth = 0 + eff_shape = (height + pad_height * 2, width + pad_width * 2) + eff_channels = channels + eff_depth = 0 + noise_dims = 2 + else: + eff_channels = channels if not depth_over_channels else 1 + eff_depth = depth if not depth_over_channels else depth * channels + eff_shape = ( + eff_depth + pad_depth * 2, + height + pad_height * 2, + width + pad_width * 2, + ) + noise_dims = 3 + + bn = base_noise + if bn is not None: + if depth_over_channels and depth > 0: + # Flattens 5D base_noise into a contiguous sequential depth format matching outputs + # (B, C, D, H, W) -> (B, 1, D*C, H, W) + bn = bn.movedim(1, 2).reshape(batch_size, 1, eff_depth, height, width) + + if pad_width > 0 or pad_height > 0 or pad_depth > 0: + if noise_dims == 3: + pad_tuple = ( + pad_width, + pad_width, + pad_height, + pad_height, + pad_depth, + pad_depth, + ) + else: + pad_tuple = (pad_width, pad_width, pad_height, pad_height) + bn = torch.nn.functional.pad(bn, pad_tuple, mode=self.pad_mode) + + noise_values = self.generate_octaves( + eff_shape, + batch_size=batch_size, + channels=eff_channels, + base_noise=bn, + ) + + # Apply normalization to the spatial dimensions individually per batch and channel + norm_dims = tuple(range(-len(eff_shape), 0)) + noise_values = normalize_to_scale(noise_values, -1.0, 1.0, dim=norm_dims) + + if self.break_pattern != 0.0: + result = self.pattern_break_blend( + noise_values, + self.break_pattern_func( + noise_values, + detail=self.detail_level, + use_frac=self.break_pattern_use_frac, + multiplier=self.break_pattern_multiplier, + ), + self.break_pattern, + ) + else: + result = noise_values + + if sum(self.pad[:3]) > 0: + if noise_dims == 3: + result = result[ + :, + :, + pad_depth : eff_depth + pad_depth, + pad_height : height + pad_height, + pad_width : width + pad_width, + ] + else: + result = result[ + :, + :, + pad_height : height + pad_height, + pad_width : width + pad_width, + ] + + if depth_over_channels and depth > 0: + # Map the contiguous sequential flattened depth back identically matching (D0_C0, D0_C1, ...) interleaving + # (B, 1, D*C, H, W) -> (B, C, D, H, W) + result = result.reshape( + batch_size, + depth, + channels, + height, + width, + ).movedim(1, 2) + + if noise_dims == 3: + # Shift Depth index backwards ahead of Batch (per expectation of the make_noise_sampler) + # (B, C, D, H, W) -> (D, B, C, H, W) + result = result.movedim(2, 0) + + return result.contiguous() + + +class PerlinItem(CustomNoiseItemBase): + def __init__( + self, + factor, + *, + perlin: Perlin | None = None, + device=None, + normalized=None, + base_noise_opt=None, + **kwargs: Any, + ): + if perlin is None: + perlin = Perlin.build(**kwargs) + super().__init__( + factor, + perlin=perlin, + device=device, + normalized=normalized + if not isinstance(normalized, str) + else NormalizeNoiseNodeMixin.get_normalize(normalized), + base_noise_opt=base_noise_opt.clone() + if base_noise_opt is not None + else None, + # **kwargs, + ) + def make_noise_sampler( self, x: torch.Tensor, @@ -426,68 +572,81 @@ class PerlinItem(CustomNoiseItemBase): seed: int | None, cpu: bool = True, normalized=True, - ): + ) -> torch.Tensor: # ty:ignore[invalid-method-override] normalized = self.get_normalize("normalized", normalized) - cpu = cpu if self.device == "default" else self.device == "cpu" - device = "cpu" if cpu else model_management.get_torch_device() + cpu = cpu if self.device == "default" and cpu else self.device == "cpu" + device = torch.device("cpu") if cpu else model_management.get_torch_device() + perlin: Perlin = self.perlin._replace(device=device, dtype=x.dtype) noise_chunk = None - noise_index = self.initial_depth + noise_index = perlin.initial_depth max_idx = None if x.ndim < 4: raise ValueError("Can only handle latents with 4+ dimensions") orig_shape = x.shape - b = x.shape[0] - c = math.prod(x.shape[1:-2]) # Hack to deal with video models. - h, w = x.shape[-2:] + b = orig_shape[0] + c = math.prod(orig_shape[1:-2]) # Hack to deal with video models + h, w = orig_shape[-2:] + + base_noise_sampler = ( + self.base_noise_opt.make_noise_sampler( + x, + sigma_min, + sigma_max, + seed=seed, + cpu=cpu, + normalized=False, + ) + if self.base_noise_opt + else None + ) + x_device, x_dtype = x.device, x.dtype del x - blend = filtering.BLENDING_MODES[self.blend] - pattern_break_blend = filtering.BLENDING_MODES[self.pattern_break_blend] - def noise_sampler(_s, _sn): + def noise_sampler(s, sn): nonlocal noise_chunk, noise_index, max_idx if noise_chunk is None: - # print("-->", noise_index, self.depth) - noise_chunk = create_noisy_latents_perlin( + base_noise = None + if base_noise_sampler: + bn_tuple = tuple( + base_noise_sampler(s, sn).reshape(b, c, h, w) + for _ in range(max(1, perlin.depth)) + ) + base_noise = ( + bn_tuple[0] + if perlin.depth < 1 + else torch.stack(bn_tuple, dim=0).movedim(0, 2) + ) + del bn_tuple + + noise_chunk = perlin( w, h, - self.depth, batch_size=b, channels=c, - detail_level=self.detail_level, - octaves=self.octaves, - persistence=self.persistence, - lacunarity=self.lacunarity, - initial_frequency=self.initial_frequency, - initial_amplitude=self.initial_amplitude, - break_pattern=self.break_pattern, - res=self.res, - tileable=self.tileable, - blend=blend, - pattern_break_blend=pattern_break_blend, - depth_over_channels=self.depth_over_channels, - pad=self.pad, - device=device, + base_noise=base_noise, ).to(device=x_device, dtype=x_dtype) - if self.depth < 1: # 2D mode + + if perlin.depth < 1: noise = noise_chunk noise_chunk = None return scale_noise(noise, self.factor, normalized=normalized) - if self.max_depth != 0 and self.max_depth != -1: - noise_chunk = noise_chunk[: self.max_depth] + if perlin.max_depth != 0 and perlin.max_depth != -1: + noise_chunk = noise_chunk[: perlin.max_depth] chunk_shape = noise_chunk.shape max_idx = ( chunk_shape[0] - 1 - if self.wrap_depth == 0 - else min(self.wrap_depth, chunk_shape[0] - 1) + if perlin.wrap_depth == 0 + else min(perlin.wrap_depth, chunk_shape[0] - 1) ) if max_idx < 0: max_idx += chunk_shape[0] + noise = noise_chunk[noise_index] noise_index += 1 if noise_index > max_idx: noise_index = 0 - if not self.wrap_depth: + if not perlin.wrap_depth: noise_chunk = None result = scale_noise(noise, self.factor, normalized=normalized) return result.reshape(orig_shape) if result.shape != orig_shape else result diff --git a/py/expression/__init__.py b/py/expression/__init__.py index 2ef5a04..d813fbd 100644 --- a/py/expression/__init__.py +++ b/py/expression/__init__.py @@ -1,8 +1,12 @@ -from . import types, expression, handler, util, validation - +from . import expression, types, util from .expression import Expression -from .validation import Arg, ValidateArg -from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext + +try: + from . import handler, validation + from .handler import BASIC_HANDLERS, BaseHandler, HandlerContext + from .validation import Arg, ValidateArg +except (ImportError, ModuleNotFoundError): + pass __all__ = ( "Arg", diff --git a/py/expression/expression.py b/py/expression/expression.py index d4ff4ae..474563a 100644 --- a/py/expression/expression.py +++ b/py/expression/expression.py @@ -1,20 +1,21 @@ -import re import operator +import re from tqdm import tqdm -from .parser import Parser, ParserSpec, ParseError +from .parser import ParseError, Parser, ParserSpec from .types import ( Empty, ExpBase, - ExpOp, ExpBinOp, - ExpSym, - ExpStatements, - ExpFunAp, - ExpTuple, ExpDict, + ExpFunAp, ExpKV, + ExpMethodAp, + ExpOp, + ExpStatements, + ExpSym, + ExpTuple, ) COMMA_PRECEDENCE = 2 @@ -36,10 +37,11 @@ class Expression: | :> # Key value binop | := # Assignment | ; # Sequencing + | :: # Method call | [?:] # Ternary | \[ | ] # Index | \.\.\. # Index ellipsis - | '[-\w.]+ # Symbol + | '[-\w.:=]+ # Symbol | `?[a-z][\w.]*`? # Function/variable names ) \s* @@ -157,9 +159,9 @@ class ExprParserSpec(ParserSpec): def split_funap_args(toks): if not isinstance(toks, (list, tuple)): return ExpTuple((toks,)), ExpDict() - return ExpTuple(t for t in toks if not isinstance(t, ExpKV)), ExpDict({ - str(t.k): t.v for t in toks if isinstance(t, ExpKV) - }) + return ExpTuple(t for t in toks if not isinstance(t, ExpKV)), ExpDict( + {str(t.k): t.v for t in toks if isinstance(t, ExpKV)} + ) @staticmethod def null_constant(p, token, bp): @@ -200,6 +202,14 @@ class ExprParserSpec(ParserSpec): p.expect(")") return make_funap(left, *cls.split_funap_args(args)) + @classmethod + def left_methodcall(cls, p, token, left, bp): + methname = p.parse_until(31) + p.expect("(") + funap = cls.left_funcall(p, token=None, left=methname, bp=None) + funap.args = ExpTuple((Empty, *funap.args)) + return ExpMethodAp(left, funap) + @staticmethod def left_comma(p, token, left, bp): if p.token == ")": @@ -251,6 +261,7 @@ class ExprParserSpec(ParserSpec): def populate(self): self.add_left(31, self.left_funcall, ("(",)) self.add_left(31, self.left_index, ("[",)) + self.add_left(31, self.left_methodcall, ("::",)) self.add_leftright(29, self.left_binop, ("**",)) self.add_null(27, self.null_prefixop, ("+", "-", "!")) self.add_left(25, self.left_binop, ("*", "/")) diff --git a/py/expression/handler.py b/py/expression/handler.py index b2aac1c..e0287f8 100644 --- a/py/expression/handler.py +++ b/py/expression/handler.py @@ -1,8 +1,9 @@ import operator +import traceback -from .validation import ValidateArg, Arg, ValidateError -from .types import Empty, ExpDict, ExpOp +from .types import Empty, ExpDict, ExpOp, ExpReturn from .util import torch +from .validation import Arg, ValidateArg, ValidateError class HandlerError(Exception): @@ -63,7 +64,8 @@ class BaseHandler: val = self.handle(obj, getter) return self.validate_output(obj, val) except Exception as exc: - raise HandlerError(f'Error evaluating "{obj.name}":\n {exc!r}') from exc + tb = traceback.format_exc() + raise HandlerError(f'Error evaluating "{obj.name}": {exc!s}\n{tb}') from exc def safe_get(self, key, obj, getter=None, *, default=Empty): str_key = isinstance(key, str) @@ -365,6 +367,13 @@ class SetVarHandler(BaseHandler): return val +class ReturnHandler(BaseHandler): + input_validators = (Arg.present("expression"),) + + def handle(self, obj, getter): + raise ExpReturn(self.safe_get("expression", obj, getter)) + + LOGIC_HANDLERS = { "||": OrHandler(), "&&": AndHandler(), @@ -400,6 +409,9 @@ MATH_HANDLERS = { ">=": RelComparisonHandler(operator.ge), "min": MinHandler(), "max": MaxHandler(), + "float": UnarySimpleMathHandler(handler=float), + "int": UnarySimpleMathHandler(handler=int), + "bool": UnarySimpleMathHandler(handler=bool), } for k, alias in ( ("+", "add"), @@ -420,6 +432,7 @@ MISC_HANDLERS = { "dict": DictHandler(), "comment": CommentHandler(), "set_var": SetVarHandler(), + "return": ReturnHandler(), } BASIC_HANDLERS = LOGIC_HANDLERS | MATH_HANDLERS | MISC_HANDLERS diff --git a/py/expression/types.py b/py/expression/types.py index bae45be..60c0f9f 100644 --- a/py/expression/types.py +++ b/py/expression/types.py @@ -2,6 +2,8 @@ class Empty: def __bool__(self): return False +class ExpReturn(Exception): + pass class ExpBase: def __bool__(self): @@ -80,8 +82,6 @@ class ExpDict(dict, ExpBase): def clone(self): return self.__class__(v.clone() if isinstance(ExpBase) else v for v in self) - def pop(self, *args, **kwargs): - raise NotImplementedError def get_eval(self, k, handlers, *args, default=Empty, **kwargs): val = super().get(k, default) @@ -106,12 +106,17 @@ class ExpDict(dict, ExpBase): for k, v in self.items() } - popitem = pop - update = pop - clear = pop - __delitem__ = pop - __setitem__ = pop - __ior__ = pop + # Can't remember if there was a compelling reason ExpDict can't be mutable but + # it breaks deep copy stuff. + # + # def pop(self, *args, **kwargs): + # raise NotImplementedError + # popitem = pop + # update = pop + # clear = pop + # __delitem__ = pop + # __setitem__ = pop + # __ior__ = pop class ExpStatements(ExpBase): @@ -135,24 +140,78 @@ class ExpStatements(ExpBase): class ExprGetter: - def __init__(self, obj, ctx, *args, **kwargs): + def __init__(self, obj, ctx, args, kwargs, *, prepend_args=()): self.obj = obj self.ctx = ctx self.args = args + self.prepend_args = prepend_args self.kwargs = kwargs def __call__(self, k, *, default=Empty): obj = self.obj - result = ( - obj.kwargs.get_eval(k, self.ctx, *self.args, default=default, **self.kwargs) - if isinstance(k, str) - else obj.args.get_eval(k, self.ctx, *self.args, **self.kwargs) - ) + if isinstance(k, str): + result = obj.kwargs.get_eval( + k, self.ctx, *self.args, default=default, **self.kwargs + ) + elif isinstance(k, int): + pa = self.prepend_args + pa_len = len(pa) + result = ( + pa[k] + if k < pa_len + else obj.args.get_eval(k, self.ctx, *self.args, **self.kwargs) + ) if result is Empty: raise KeyError(f"Unknown key {k!r}") return result +class ExpMethodAp(ExpBase): + __slots__ = ("object_expression", "funap") + + def __init__(self, object_expression, funap): + super().__init__() + self.object_expression = object_expression + self.funap = funap + + def eval(self, handlers, *args, **kwargs): + object_value = self.object_expression.eval(handlers, *args, **kwargs) + type_name = type(object_value).__name__ + handler_key = f"{type_name}::{self.funap.name}" + handler = handlers.get_handler(handler_key) + if handler is Empty: + raise KeyError(f"No handler for method call op: {handler_key!r}") + return handler( + self, + getter=ExprGetter( + self.funap, handlers, args, kwargs, prepend_args=(object_value,) + ), + **kwargs, + ) + + def clone(self): + return self.__class__( + object_expression=self.object_expression.clone(), + funap=self.funap.clone(), + ) + __copy__ = clone + + def __getattr__(self, k): + if k == "name": + return f"method::{self.funap.name}" + if k == "args": + return self.funap.args + if k == "kwargs": + return self.funap.kwargs + # This doesn't play well with deep copy. + # if hasattr(self.funap, k): + # return getattr(self.funap, k) + raise AttributeError(f"Can't get attribute {k}") + + def __repr__(self): + return f"" + + class ExpFunAp(ExpBase): __slots__ = ("name", "args", "kwargs") @@ -165,9 +224,7 @@ class ExpFunAp(ExpBase): handler = handlers.get_handler(self.name) if handler is Empty: raise KeyError(f"No handler for op: {self.name!r}") - return handler( - self, getter=ExprGetter(self, handlers, *args, **kwargs), **kwargs - ) + return handler(self, getter=ExprGetter(self, handlers, args, kwargs), **kwargs) def clone(self): return self.__class__(self.name, self.args.clone(), self.kwargs.clone()) @@ -209,5 +266,6 @@ __all__ = ( "ExpKV", "ExpDict", "ExpFunAp", + "ExpMethodAp", "ExpBoundFunAp", ) diff --git a/py/expression/validation.py b/py/expression/validation.py index 1a089f6..daaeec2 100644 --- a/py/expression/validation.py +++ b/py/expression/validation.py @@ -2,8 +2,8 @@ import contextlib import functools from ..latent import ImageBatch -from .util import torch from .types import Empty +from .util import torch class Arg: @@ -53,6 +53,17 @@ class Arg: name, default=default, validator=ValidateArg.validate_numscalar_sequence ) + @classmethod + def numscalar_sequence_or_single(cls, name, default=Empty): + return cls.one_of( + name, + ( + ValidateArg.validate_numscalar_sequence, + ValidateArg.validate_numeric_scalar, + ), + default=default, + ) + @classmethod def tensor_slice(cls, name, default=Empty): return cls(name, default=default, validator=ValidateArg.validate_tensor_slice) @@ -236,6 +247,12 @@ class ValidateArg: raise ValidateError(f"Expected string argument at {idx}, got {type(val)}") return val + @classmethod + def validate_dict(cls, idx, val): + if not isinstance(val, dict): + raise ValidateError(f"Expected dict argument at {idx}, got {type(val)}") + return val + @classmethod def validate_boolean(cls, idx, val): if val is not True and val is not False: diff --git a/py/expression_handlers.py b/py/expression_handlers.py index 2910a4e..6cec4fc 100644 --- a/py/expression_handlers.py +++ b/py/expression_handlers.py @@ -1,18 +1,15 @@ import os - -import torch -import numpy as np -import PIL.Image as PILImage from functools import partial +import numpy as np +import PIL.Image as PILImage +import torch + from . import expression as expr -from . import latent -from . import unsafe_expression_whitelists - +from . import latent, unsafe_expression_whitelists from .external import MODULES as EXT -from .utils import scale_noise, resolve_value, quantile_normalize from .latent import OCSTAESD, ImageBatch, normalize_to_scale - +from .utils import quantile_normalize, resolve_value, scale_noise ALLOW_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_UNSAFE_EXPRESSIONS") is not None ALLOW_ALL_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_ALL_UNSAFE") is not None @@ -107,6 +104,14 @@ class ClampHandler(NormHandler): return torch.clamp(tensor, min=tmin, max=tmax) +class AbsHandler(NormHandler): + input_validators = (expr.Arg.tensor("tensor"),) + + def handle(self, obj, getter): + (tensor,) = self.safe_get_all(obj, getter) + return tensor.abs() + + class StackHandler(NormHandler): input_validators = ( expr.Arg.sequence("tensors", item_validator=expr.ValidateArg.validate_tensor), @@ -135,17 +140,52 @@ class ReshapeHandler(NormHandler): return torch.reshape(tensor.clone(), shape) +class SplitHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.integer("chunk_size"), + expr.Arg.integer("dim"), + expr.Arg.boolean("pad_last", False), + ) + + def handle(self, obj, getter): + tensor, chunk_size, dim, pad_last = self.safe_get_all(obj, getter) + chunks = torch.split(tensor, chunk_size, dim=dim) + if not chunks or not pad_last or chunks[-1].shape == chunks[0].shape: + return chunks + dsize = chunks[0].shape[dim] + replacement_chunk = tensor.new_zeros(chunks[0].shape) + replacement_chunk[ + tuple( + slice(None, dsize if d == dim else None) for d in replacement_chunk.ndim + ) + ] = chunks[-1] + return (*chunks[:-1], replacement_chunk) + + +class TrimHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.numscalar_sequence("shape"), + ) + + def handle(self, obj, getter): + tensor, shape = self.safe_get_all(obj, getter) + return tensor[tuple(slice(0, dsize) for dsize in shape)] + + class IndexedCopyHandler(NormHandler): input_validators = ( expr.Arg.tensor("tensor_dest"), expr.Arg.tensor("tensor_src"), expr.Arg.tensor_slice("slice"), + expr.Arg.boolean("slice_src", default=True), ) def handle(self, obj, getter): - tensor1, tensor2, tensor_slice = self.safe_get_all(obj, getter) + tensor1, tensor2, tensor_slice, slice_src = self.safe_get_all(obj, getter) result = tensor1.clone() - result[tensor_slice] = tensor2[tensor_slice] + result[tensor_slice] = tensor2[tensor_slice] if slice_src else tensor2 return result @@ -167,22 +207,24 @@ class NewLikeHandler(NormHandler): class MeanHandler(NormHandler): input_validators = ( expr.Arg.tensor("tensor"), - expr.Arg.numscalar_sequence("dim", (-3, -2, -1)), + expr.Arg.numscalar_sequence_or_single("dim", (-3, -2, -1)), ) def handle(self, obj, getter): tensor, dim = self.safe_get_all(obj, getter) + dim = dim if isinstance(dim, tuple) else (dim,) return tensor.mean(keepdim=True, dim=dim) class StdHandler(NormHandler): input_validators = ( expr.Arg.tensor("tensor"), - expr.Arg.numscalar_sequence("dim", (-3, -2, -1)), + expr.Arg.numscalar_sequence_or_single("dim", (-3, -2, -1)), ) def handle(self, obj, getter): tensor, dim = self.safe_get_all(obj, getter) + dim = dim if isinstance(dim, tuple) else (dim,) return tensor.std(keepdim=True, dim=dim) @@ -190,17 +232,7 @@ class RollHandler(NormHandler): input_validators = ( expr.Arg.tensor("tensor"), expr.Arg.numeric_scalar("amount", 0.5), - expr.Arg.one_of( - "dim", - ( - expr.ValidateArg.validate_integer, - partial( - expr.ValidateArg.validate_sequence, - item_validator=expr.ValidateArg.validate_integer, - ), - ), - default=-2, - ), + expr.Arg.numscalar_sequence_or_single("dim", -2), ) def handle(self, obj, getter): @@ -280,15 +312,34 @@ class BlendHandler(NormHandler): expr.Arg.tensor("tensor1"), expr.Arg.tensor("tensor2"), expr.Arg.numeric("scale", 0.5), - expr.Arg.string("mode", "lerp"), + expr.Arg.one_of( + "mode", + ( + expr.ValidateArg.validate_string, + expr.ValidateArg.validate_dict, + ), + ), + expr.Arg.one_of( + "blend_kwargs", + ( + expr.ValidateArg.validate_none, + expr.ValidateArg.validate_dict, + ), + default=None, + ), ) def handle(self, obj, getter): - t1, t2, scale, mode = self.safe_get_all(obj, getter) + t1, t2, scale, mode, blend_kwargs = self.safe_get_all(obj, getter) blend_handler = BLENDING_MODES.get(mode) if not blend_handler: raise KeyError(f"Unknown blend mode {mode!r}") - return blend_handler(t1, t2, scale) + return blend_handler( + t1, + t2, + scale, + **({} if blend_kwargs is None else blend_kwargs), + ) class ContrastAdaptiveSharpeningHandler(NormHandler): @@ -633,8 +684,11 @@ TENSOR_OP_HANDLERS = { "t_normalize_to_scale": NormToScaleHandler(), "t_reshape": ReshapeHandler(), "t_clamp": ClampHandler(), + "t_abs": AbsHandler(), "t_cat": CatHandler(), "t_stack": StackHandler(), + "t_split": SplitHandler(), + "t_trim": TrimHandler(), "t_indexed_copy": IndexedCopyHandler(), "t_new_like": NewLikeHandler(), "t_mean": MeanHandler(), @@ -657,6 +711,11 @@ TENSOR_OP_HANDLERS = { "unsafe_torch": UnsafeTorchHandler(), } +TENSOR_OP_HANDLERS |= { + f"Tensor::{k[2:] if k.startswith('t_') else k}": v + for k, v in TENSOR_OP_HANDLERS.items() +} + IMAGE_OP_HANDLERS = { "img_taesd_encode": TAESDEncodeHandler(), "img_shape": ImgShapeHandler(), diff --git a/py/model.py b/py/model.py index 43cb23e..8f21232 100644 --- a/py/model.py +++ b/py/model.py @@ -169,8 +169,9 @@ class OCSModel: self.extra_args = extra_args self.cfg1_uncond_optimization = cfg1_uncond_optimization self.cfg_scale_override = cfg_scale_override + self.model_sampling = model.inner_model.inner_model.model_sampling self.is_rectified_flow = isinstance( - model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST + self.model_sampling, comfy.model_sampling.CONST ) self.latent_format = OCSLatentFormat( x.device, model.inner_model.inner_model.latent_format @@ -212,9 +213,9 @@ class OCSModel: ) -> torch.Tensor: return self.model(x, sigma * self.s_in, **self.extra_args | kwargs) - @property - def model_sampling(self): - return self.model.inner_model.inner_model.model_sampling + # @property + # def model_sampling(self): + # return self.model.inner_model.inner_model.model_sampling @property def inner_cfg_scale(self) -> None | int | float: diff --git a/py/nodes.py b/py/nodes.py index d82023b..c1ab792 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -1,17 +1,17 @@ +from __future__ import annotations + import comfy -import yaml - import torch - +import yaml from tqdm import tqdm from .external import MODULES, IntegratedNode +from .filtering import Filter, FilterRefs, make_filter from .restart import Restart from .sampling import composable_sampler from .step_samplers import STEP_SAMPLERS from .substep_merging import MERGE_SUBSTEPS_CLASSES from .substep_sampling import ParamGroup, StepSamplerChain, StepSamplerGroups -from .filtering import make_filter try: from comfy_execution import validation as comfy_validation @@ -27,15 +27,17 @@ except Exception as exc: f"** OCS: Warning, caught unexpected exception trying to detect ComfyUI union type support. Disabling. Exception: {exc}" ) -PARAM_INPUT_TYPES = frozenset(( - "IMAGE", - "OCS_NOISE", - "SAMPLER", - "SIGMAS", - "SONAR_CUSTOM_NOISE", - "UPSCALE_MODEL", - "VAE", -)) +PARAM_INPUT_TYPES = frozenset( + ( + "IMAGE", + "OCS_NOISE", + "SAMPLER", + "SIGMAS", + "SONAR_CUSTOM_NOISE", + "UPSCALE_MODEL", + "VAE", + ) +) NOISE_INPUT_TYPES = frozenset(("SONAR_CUSTOM_NOISE", "OCS_NOISE")) @@ -778,6 +780,329 @@ class ApplyFilterImage(ApplyFilterLatent): return (result,) +class ExpressionFilteredLatentOperation: + EXTENDED_LATENT_OPERATION = True + + def __init__( + self, + *, + ocs_filter: Filter, + latent_refs: dict[str, torch.Tensor] | None = None, + ) -> None: + self.filter = ocs_filter + self.latent_refs = latent_refs if latent_refs is not None else {} + + def __call__( + self, + latent: torch.Tensor, + *, + sigma: float | torch.Tensor | None = None, + **kwargs: dict, + ) -> torch.Tensor: + refs = FilterRefs( + kvs={ + "sigma": sigma.clone() if isinstance(sigma, torch.Tensor) else sigma, + "sigma_float": sigma.max().item() + if isinstance(sigma, torch.Tensor) + else sigma, + } + | {k: v.to(latent, copy=True) for k, v in self.latent_refs.items()} + ) + return self.filter.apply(latent, refs=refs) + + +class ExpressionFilteredLatentOperationNode: + DESCRIPTION = "TBD" + + FUNCTION = "go" + RETURN_TYPES = ("LATENT_OPERATION",) + + @classmethod + def INPUT_TYPES(cls): + MODULES.initialize() + return { + "required": { + "yaml_config": ( + "STRING", + { + "default": "", + "placeholder": """\ + # YAML or JSON filter definition + """, + "multiline": True, + "dynamicPrompts": False, + "tooltip": "Enter your filter definition here. There is essentially no error handling.", + }, + ), + }, + "optional": { + "latent_ref_1_opt": ("LATENT",), + "latent_ref_2_opt": ("LATENT",), + "latent_ref_3_opt": ("LATENT",), + }, + } + + def go( + self, + *, + yaml_config: str, + latent_ref_1_opt: dict | None = None, + latent_ref_2_opt: dict | None = None, + latent_ref_3_opt: dict | None = None, + ) -> tuple: + config = yaml.safe_load(yaml_config) + if not isinstance(config, dict) or "filter" not in config: + raise ValueError( + "Bad YAML config type (must be object) or missing filter key in config" + ) + filter_def = config.get("filter") + if not isinstance(filter_def, dict): + raise ValueError("Bad type for filter definition, must be object") + latent_refs = { + k: v["samples"].to(device="cpu", dtype=torch.float32, copy=True) + for k, v in ( + ("latent_ref_1", latent_ref_1_opt), + ("latent_ref_2", latent_ref_2_opt), + ("latent_ref_3", latent_ref_3_opt), + ) + if v is not None + } + ocs_filter = make_filter(filter_def) + return ( + ExpressionFilteredLatentOperation( + ocs_filter=ocs_filter, latent_refs=latent_refs + ), + ) + + +class ExpressionFilteredModelPatchNode: + DESCRIPTION = "TBD" + + FUNCTION = "go" + RETURN_TYPES = ("MODEL",) + + @classmethod + def INPUT_TYPES(cls): + MODULES.initialize() + return { + "required": { + "model": ("MODEL",), + "patch_mode": ( + ("apply_model", "pre_cfg", "post_cfg", "cfg", "denoise_mask"), + {"default": "apply_model"}, + ), + "existing_patch_mode": ( + ("normal", "extract", "extract_split", "extract_sequence"), + { + "default": "normal", + "tooltip": "Modes:\n" + "normal: Replaces apply_model or cfg patches, appends for pre_cfg and post_cfg.\n" + "extract: Removes the existing patches and passes old_result with the output from existing patches.\n" + "extract_split: Same as extract except you'll get tuple of results for each existing patch (apply_model and cfg will always be length 1).\n" + "extract_sequence: Like extract_split except existing patches do not see each other's effects.", + }, + ), + "yaml_config": ( + "STRING", + { + "default": "", + "placeholder": """\ + # YAML or JSON filter definition + """, + "multiline": True, + "dynamicPrompts": False, + "tooltip": "Enter your filter definition here. There is essentially no error handling.", + }, + ), + }, + "optional": { + "latent_ref_1_opt": ("LATENT",), + "latent_ref_2_opt": ("LATENT",), + "latent_ref_3_opt": ("LATENT",), + }, + } + + def go( + self, + *, + yaml_config: str, + model: object, + patch_mode: str, + existing_patch_mode: str, + latent_ref_1_opt: dict | None = None, + latent_ref_2_opt: dict | None = None, + latent_ref_3_opt: dict | None = None, + ) -> tuple: + config = yaml.safe_load(yaml_config) + if not isinstance(config, dict) or "filter" not in config: + raise ValueError( + "Bad YAML config type (must be object) or missing filter key in config" + ) + filter_def = config.get("filter") + if not isinstance(filter_def, dict): + raise ValueError("Bad type for filter definition, must be object") + latent_refs = { + k: v["samples"].to(device="cpu", dtype=torch.float32, copy=True) + for k, v in ( + ("latent_ref_1", latent_ref_1_opt), + ("latent_ref_2", latent_ref_2_opt), + ("latent_ref_3", latent_ref_3_opt), + ) + if v is not None + } + ocs_filter = make_filter(filter_def) + model = model.clone() + mode_keys = { + "post_cfg": "sampler_post_cfg_function", + "pre_cfg": "sampler_pre_cfg_function", + "apply_model": "model_function_wrapper", + "cfg": "sampler_cfg_function", + "denoise_mask": "denoise_mask_function", + } + key = mode_keys.get(patch_mode) + if key is None: + raise ValueError(f"Bad mode: {patch_mode}") + if existing_patch_mode != "normal": + old_handlers = model.model_options.pop(key, None) + if old_handlers is None: + old_handlers = () + else: + old_handlers = () + + def get_refs(*args, **kwargs) -> FilterRefs: + old_results = [] + if patch_mode in {"pre_cfg", "post_cfg", "cfg"}: + argdict = args[0] + elif patch_mode == "denoise_mask": + argdict = { + "sigma": args[0], + "denoise_mask": args[1].clone(), + "sigmas": kwargs["extra_options"]["sigmas"].clone(), + } + elif patch_mode == "apply_model": + argdict = args[1] | {"apply_function": args[0]} + else: + raise ValueError(f"Bad patch mode: {patch_mode}") + if old_handlers: + ridx = 0 if existing_patch_mode == "extract_sequence" else -1 + if patch_mode in {"cfg", "denoise_mask", "apply_model"}: + old_results = (old_handlers[0](*args, **kwargs),) + elif patch_mode == "pre_cfg": + old_results = [argdict["conds_out"]] + for hf in old_handlers: + result = hf(argdict | {"conds_out": old_results[ridx]}).copy() + if len(old_results) > 1 and existing_patch_mode != "extract": + old_results[1] = result + else: + old_results.append(result) + old_results = old_results[1:] + elif patch_mode == "post_cfg": + old_results = [argdict["denoised"].clone()] + for hf in old_handlers: + result = hf(argdict | {"denoised": old_results[ridx].clone()}) + if len(old_results) > 1 and existing_patch_mode != "extract": + old_results[1] = result + else: + old_results.append(result) + old_results = old_results[1:] + kvs = { + "sigma": argdict["sigma"].clone(), + "sigma_float": argdict["sigma"].max().item(), + "old_results": tuple(old_results), + } + if patch_mode in {"pre_cfg", "post_cfg", "cfg"}: + kvs |= { + "x": argdict["input"].clone(), + "cfg_scale": argdict["cond_scale"], + } + if patch_mode in {"post_cfg", "cfg"}: + kvs["cond"] = argdict["cond_denoised"].clone() + uncond = argdict.get("uncond_denoised", None) + kvs["uncond"] = uncond if uncond is None else uncond.clone() + if patch_mode == "post_cfg": + kvs["denoised"] = argdict["denoised"].clone() + else: + conds_out = argdict["conds_out"] + kvs["cond"] = conds_out[0].clone() + kvs["uncond"] = ( + conds_out[1].clone() + if len(conds_out) > 1 and conds_out[1] is not None + else None + ) + kvs["conds_out"] = list(conds_out) + elif patch_mode == "denoise_mask": + kvs |= { + "sigmas": argdict["sigmas"], + "denoise_mask": argdict["denoise_mask"], + } + elif patch_mode == "apply_model": + kvs |= { + "x": argdict["input"].clone(), + "cond_or_uncond": argdict["cond_or_uncond"].clone(), + } + else: + raise ValueError(f"Bad patch mode: {patch_mode}") + if patch_mode in {"pre_cfg", "post_cfg", "cfg", "apply_model"}: + latent_in = kvs["x"] + else: + latent_in = kvs["denoise_mask"] + kvs |= {k: v.to(latent_in) for k, v in latent_refs.items()} + return FilterRefs(kvs=kvs) + + def model_patch(*args, **kwargs): + refs = get_refs(*args, **kwargs) + if patch_mode == "apply_model": + + def fallback_apply_model(): + old_results = refs.kvs["old_results"] + if old_results: + return old_results[-1] + return args[0]( + args[1]["input"], args[1]["timestep"], **args[1]["c"] + ) + else: + fallback_apply_model = None + if not ocs_filter.check_applies(refs): + old_results = refs.kvs["old_results"] + if old_results: + return old_results[-1] + if patch_mode == "pre_cfg": + return args[0]["conds_out"] + if patch_mode == "post_cfg": + return args[0]["denoised"] + if patch_mode == "cfg": + return args[0]["cond"] + if patch_mode == "denoise_mask": + return args[1] + if patch_mode == "apply_model": + return fallback_apply_model() + raise ValueError(f"Bad patch mode: {patch_mode}") + if patch_mode in {"pre_cfg", "post_cfg", "cfg", "apply_model"}: + latent_in = refs.kvs["x"] + else: + latent_in = refs.kvs["denoise_mask"] + result = ocs_filter.apply(latent_in, refs=refs) + if patch_mode == "apply_model" and result is None: + return fallback_apply_model() + if patch_mode == "pre_cfg": + return list(result) + return result + + if patch_mode == "pre_cfg": + model.set_model_sampler_pre_cfg_function(model_patch) + elif patch_mode == "post_cfg": + model.set_model_sampler_post_cfg_function(model_patch) + elif patch_mode == "cfg": + model.set_model_sampler_cfg_function(model_patch) + elif patch_mode == "denoise_mask": + model.set_model_denoise_mask_function(model_patch) + elif patch_mode == "apply_model": + model.set_model_unet_function_wrapper(model_patch) + else: + raise ValueError(f"Bad patch mode: {patch_mode}") + return (model,) + + __all__ = ( "SamplerNode", "GroupNode", diff --git a/py/noise.py b/py/noise.py index fc29d17..421b080 100644 --- a/py/noise.py +++ b/py/noise.py @@ -2,12 +2,67 @@ import gc import math import random +import numpy as np import scipy import torch from tqdm import tqdm from .filtering import Filter, make_filter -from .utils import scale_noise, fallback +from .utils import fallback, scale_noise + +try: + from .triton_lsa import ( + assignments_to_indices, + batch_linear_assignment, + batch_linear_assignment_shuffled, + ) + + HAVE_TRITON = True +except Exception: + HAVE_TRITON = False + + +def linear_sum_assignment( + cost: torch.Tensor, + *, + maximize: bool = False, + use_triton: bool = False, + split_batch: int = 0, + **kwargs: dict, +) -> tuple[np.ndarray, np.ndarray] | tuple[torch.Tensor, torch.Tensor]: + if not use_triton or not HAVE_TRITON or not cost.is_cuda: + cost = cost.half().cpu() + return scipy.optimize.linear_sum_assignment(cost, maximize=maximize) + ndim = cost.ndim + orig_shape = cost.shape + if ndim == 2: + do_split = split_batch > 1 and all( + (sz / split_batch).is_integer() for sz in orig_shape + ) + if do_split: + cost = cost.reshape( + split_batch, orig_shape[0] // split_batch, orig_shape[1] // split_batch + ) + else: + cost = cost.unsqueeze(0) + tqdm.write( + f"TRITON LAP: maximize={maximize}, orig cost shape={orig_shape}, cost shape={cost.shape}, cost dtype={cost.dtype}", + ) + if not cost.is_contiguous(): + cost = cost.contiguous() + fun = ( + batch_linear_assignment + if "generator" not in kwargs + else batch_linear_assignment_shuffled + ) + assignments = fun(cost, maximize=maximize, **kwargs) + row_ind, col_ind = assignments_to_indices(assignments) + if ndim == 2: + row_ind = row_ind.reshape(-1, row_ind.shape[-1]) + col_ind = col_ind.reshape(-1, col_ind.shape[-1]) + # row_ind, col_ind = row_ind.squeeze(0), col_ind.squeeze(0) + tqdm.write(f"Ran LAP kernel: {assignments.shape}, {row_ind.shape}, {col_ind.shape}") + return row_ind, col_ind class ImmiscibleNoise(Filter): @@ -19,6 +74,12 @@ class ImmiscibleNoise(Filter): "maximize": False, "distance_scale": 0.0, "distance_scale_ref": None, + "abs_mode": False, + "abs_distance_mode": False, + "use_triton": False, + # Only honored in Triton mode. + "split_batch": 0, + "generator": None, } def __call__(self, noise_sampler, x_ref, *, refs=None): @@ -83,7 +144,11 @@ class ImmiscibleNoise(Filter): # "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303 # Minimize latent-noise pairs over a batch batch = latent.shape[0] + out_latent = fallback(out_latent, latent) ref_latent = ref_latent.detach().clone() + if self.abs_mode: + ref_latent = ref_latent.abs() + latent = latent.abs() if self.distance_scale == 0: ref_latent_expanded = ref_latent.unsqueeze(1).expand( -1, batch, *ref_latent.shape[1:] @@ -92,30 +157,30 @@ class ImmiscibleNoise(Filter): ref_latent.shape[0], *latent.shape ) dist = (ref_latent_expanded - latent_expanded) ** 2 + if self.abs_distance_mode: + dist = dist.abs_() del ref_latent_expanded, latent_expanded - dist = dist.mean(tuple(range(2, dist.dim()))) + dist = dist.mean(tuple(range(2, dist.ndim))) else: - dist = torch.linalg.vector_norm( - fallback(self.distance_scale_ref, self.distance_scale) - * ref_latent.flatten(start_dim=1).unsqueeze(1) - - self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0), - dim=2, - ) - dist = dist.half() + distance_scale_ref = fallback(self.distance_scale_ref, self.distance_scale) + dist = distance_scale_ref * ref_latent.flatten(start_dim=1).unsqueeze( + 1 + ) - self.distance_scale * latent.flatten(start_dim=1).unsqueeze(0) + if self.abs_distance_mode: + dist = dist.abs_() + dist = torch.linalg.vector_norm(dist, dim=2) try: - assign_mat = scipy.optimize.linear_sum_assignment( - dist.cpu(), maximize=self.maximize + assign_mat = linear_sum_assignment( + dist, + maximize=self.maximize, + use_triton=self.use_triton, + split_batch=self.split_batch, + generator=self.generator, ) except ValueError as exc: tqdm.write(f"OCS: Immiscible: Failed due to exception: {exc}") - return ( - None - if return_idxs - else fallback(out_latent, latent)[: ref_latent.shape[0]] - ) - return ( - assign_mat if return_idxs else fallback(out_latent, latent)[assign_mat[1]] - ) + return None if return_idxs else out_latent[: ref_latent.shape[0]] + return assign_mat if return_idxs else out_latent[assign_mat[1]] def immiscible_simple( self, diff --git a/py/restart.py b/py/restart.py index 3c59614..1fb3552 100644 --- a/py/restart.py +++ b/py/restart.py @@ -1,8 +1,57 @@ +from __future__ import annotations + +from typing import NamedTuple + import torch +# from tqdm import tqdm + + +class RestartScaleFactors(NamedTuple): + latent_scale: float + noise_scale: float + + @classmethod + def build( + cls, + sigma_from: float | torch.Tensor, + sigma_to: float | torch.Tensor, + *, + is_flow: bool, + ) -> RestartScaleFactors: + if isinstance(sigma_from, torch.Tensor): + sigma_from = sigma_from.max().item() + if isinstance(sigma_to, torch.Tensor): + sigma_to = sigma_to.max().item() + if not is_flow: + return cls( + 1.0, + max(0.0, (sigma_to**2 - sigma_from**2)) ** 0.5, + ) + alpha_from = 1.0 - sigma_from + alpha_to = 1.0 - sigma_to + if alpha_to <= 0: + latent_scale = 0.0 + noise_scale = sigma_to + else: + latent_scale = alpha_to / alpha_from + noise_scale = ( + max(0.0, (sigma_to**2) - (latent_scale * sigma_from) ** 2) ** 0.5 + ) + return cls(latent_scale, noise_scale) + class Restart: - def __init__(self, *, s_noise=1.0, custom_noise=None, immiscible=False): + def __init__( + self, + *, + s_noise=1.0, + custom_noise=None, + immiscible=False, + normalized=True, + normalize_dims: tuple[int, ...] | None = None, + is_flow=False, + ): from .noise import ImmiscibleNoise self.s_noise = s_noise @@ -10,6 +59,9 @@ class Restart: immiscible = ImmiscibleNoise(**immiscible) self.immiscible = immiscible self.custom_noise = custom_noise + self.normalized = normalized + self.normalize_dims = normalize_dims + self.is_flow = is_flow def get_noise_sampler(self, nsc): return nsc.make_caching_noise_sampler( @@ -30,23 +82,64 @@ class Restart: last_sigma = sigma return sigmas - def split_sigmas(self, sigmas): + def split_sigmas(self, sigmas: torch.Tensor): prev_seg = None while len(sigmas) > 1: seg = self.get_segment(sigmas) sigmas = sigmas[len(seg) :] if prev_seg is not None and seg[0] > prev_seg[-1]: - noise_scale = self.get_noise_scale(prev_seg[-1], seg[0]) + scale_factors = RestartScaleFactors.build( + sigma_from=prev_seg[-1], sigma_to=seg[0], is_flow=self.is_flow + ) else: - noise_scale = 0.0 + scale_factors = None prev_seg = seg - yield (noise_scale, seg) + yield (scale_factors, seg) - def get_noise_scale(self, s_min, s_max): + def get_noise_scale( + self, s_min: float | torch.Tensor, s_max: float | torch.Tensor + ) -> float: result = (s_max**2 - s_min**2) ** 0.5 if isinstance(result, torch.Tensor): - result = result.item() - return result * self.s_noise + return result.item() + return result + + def add_noise( + self, + x: torch.Tensor, + sigma_from: float, + sigma_to: float, + *, + nsc, + refs, + scale_factors: RestartScaleFactors | None = None, + in_place: bool = False, + ) -> torch.Tensor: + if self.is_flow: + sigma_from = min(1.0, max(0.0, sigma_from)) + sigma_to = min(1.0, max(0.0, sigma_to)) + if sigma_from >= sigma_to: + raise ValueError( + f"sigma_from ({sigma_from:.4f}) must be less than sigma_to ({sigma_to:.4f})" + ) + scale_factors = scale_factors or RestartScaleFactors.build( + sigma_from, sigma_to, is_flow=self.is_flow + ) + ns = self.get_noise_sampler(nsc) + sigma_empty = nsc.min_sigma * 0 + noise = nsc.scale_noise( + ns(sigma_empty + sigma_from, sigma_empty + sigma_to, refs=refs), + normalized=self.normalized, + normalize_dims=self.normalize_dims, + ) + noise *= scale_factors.noise_scale * self.s_noise + if scale_factors.latent_scale != 1.0: + x = ( + x.mul_(scale_factors.latent_scale) + if in_place + else scale_factors.latent_scale * x + ) + return noise.add_(x) def __repr__(self): return f"" @@ -77,18 +170,27 @@ class Restart: raise ValueError("Schedule jump index out of range") sched_idx = item continue + sched_frac = round(sig_idx - int(sig_idx), ndigits=5) + sig_idx = int(sig_idx if sched_frac == 0 else sig_idx + 1) if sig_idx >= siglen or sig_idx < 0: break interval, jump = item chunk = siglist[sig_idx : sig_idx + interval + 1] + if sched_frac != 0: + chunk[0] -= (chunk[0] - chunk[1]) * (1.0 - sched_frac) # print(f"{out} + {chunk}") out += chunk - sig_idx += interval + jump if jump >= 0: sig_idx += 1 + sig_idx += interval + jump sched_idx += 1 + sched_frac = round(sig_idx - int(sig_idx), ndigits=5) + sig_idx = int(sig_idx if sched_frac == 0 else sig_idx + 1) if sig_idx < siglen and sig_idx >= 0: - out += siglist[sig_idx:] + chunk = siglist[sig_idx:] + if sched_frac != 0: + chunk[0] -= (chunk[0] - chunk[1]) * (1.0 - sched_frac) + out += chunk if out[-1] > siglist[-1]: out.append(siglist[-1]) return torch.tensor(out).to(sigmas) diff --git a/py/sampling.py b/py/sampling.py index c52575a..fa46dec 100644 --- a/py/sampling.py +++ b/py/sampling.py @@ -1,13 +1,12 @@ import torch from tqdm.auto import trange - from .filtering import FILTER_HANDLERS, FilterRefs from .model import OCSModel from .noise import NoiseSamplerCache -from .substep_sampling import SamplerState -from .substep_merging import MERGE_SUBSTEPS_CLASSES from .restart import Restart +from .substep_merging import MERGE_SUBSTEPS_CLASSES +from .substep_sampling import SamplerState def find_merge_sampler(merge_samplers, ss) -> object | None: @@ -48,11 +47,6 @@ def composable_sampler( restart_custom_noise = copts.get("restart_custom_noise") if isinstance(restart_custom_noise, str): restart_custom_noise = copts.get(f"restart_custom_noise_{restart_custom_noise}") - restart = Restart( - s_noise=restart_params.get("s_noise", 1.0), - custom_noise=restart_custom_noise, - immiscible=restart_params.get("immiscible", False), - ) ss = SamplerState( OCSModel( @@ -72,6 +66,16 @@ def composable_sampler( reta=copts.get("reta", 1.0), disable_status=disable, ) + + restart = Restart( + s_noise=restart_params.get("s_noise", 1.0), + custom_noise=restart_custom_noise, + immiscible=restart_params.get("immiscible", False), + normalized=restart_params.get("normalized", True), + normalize_dims=restart_params.get("normalize_dims"), + is_flow=ss.model.is_rectified_flow, + ) + groups = copts["_groups"] merge_samplers = tuple( MERGE_SUBSTEPS_CLASSES[g.merge_method](ss, g) for g in groups.items @@ -85,17 +89,17 @@ def composable_sampler( ) ss.noise = nsc sigma_chunks = ( - tuple(restart.split_sigmas(sigmas)) if restart_enabled else ((0.0, sigmas),) + tuple(restart.split_sigmas(sigmas)) if restart_enabled else ((None, sigmas),) ) step_count = sum(len(chunk) - 1 for _noise, chunk in sigma_chunks) ss.total_steps = step_count step = 0 with trange(step_count, disable=ss.disable_status) as pbar: - for noise_scale, chunk_sigmas in sigma_chunks: - if step != 0 and noise_scale != 0: - prev_refs = FilterRefs({ - f"pre_restart_{k}": v for k, v in ss.refs.items() - }) + for chunk_idx, (scale_factors, chunk_sigmas) in enumerate(sigma_chunks): + if step != 0 and scale_factors is not None: + prev_refs = FilterRefs( + {f"pre_restart_{k}": v for k, v in ss.refs.items()} + ) ss.sigmas = chunk_sigmas ss.update(0, step=step, substep=0) if step != 0: @@ -104,14 +108,20 @@ def composable_sampler( ss.hist.reset() for ms in merge_samplers: ms.reset() - nsc.min_sigma, nsc.max_sigma = chunk_sigmas[-1], chunk_sigmas[0] - if step != 0 and noise_scale != 0: - restart_ns = restart.get_noise_sampler(nsc) - x += nsc.scale_noise( - restart_ns(nsc.min_sigma, nsc.max_sigma, refs=prev_refs | ss.refs), - noise_scale, + nsc.min_sigma, nsc.max_sigma = ( + chunk_sigmas[-1].clone(), + chunk_sigmas[0].clone(), + ) + if step != 0 and scale_factors is not None: + x = restart.add_noise( + x, + sigma_from=sigma_chunks[chunk_idx - 1][1][-1].item(), + sigma_to=chunk_sigmas[0].item(), + scale_factors=scale_factors, + nsc=nsc, + refs=prev_refs | ss.refs, + in_place=True, ) - del restart_ns del prev_refs for idx in range(len(chunk_sigmas) - 1): if idx > 0: diff --git a/py/step_samplers/blep.py b/py/step_samplers/blep.py index aa91433..ba949e4 100644 --- a/py/step_samplers/blep.py +++ b/py/step_samplers/blep.py @@ -1,20 +1,19 @@ +import inspect +import math import typing -import inspect +import comfy import torch -import comfy - -from .. import filtering from .. import expression as expr +from .. import filtering from ..utils import fallback from .base import ( - StepSamplerContext, SingleStepSampler, + StepSamplerContext, registry, ) - try: import pytorch_wavelets as ptwav @@ -531,7 +530,10 @@ class WeoonStep(SingleStepSampler): ) denoised_new = self.wavelet_inverse(coeffs_out) if denoised_new.shape != x.shape: - denoised_new = denoised_new.reshape(*x.shape) + bi_elements = math.prod(x.shape[1:]) + denoised_new = denoised_new.reshape(x.shape[0], -1)[ + :, :bi_elements + ].reshape(*x.shape) x = self.blend(denoised_new, x, ratio) yield from self.result(x, sigma_up) diff --git a/py/step_samplers/builtins.py b/py/step_samplers/builtins.py index 3ee11fa..66423fd 100644 --- a/py/step_samplers/builtins.py +++ b/py/step_samplers/builtins.py @@ -1,13 +1,14 @@ -import torch - import comfy +import torch from comfy.k_diffusion.sampling import get_ancestral_step +from tqdm import tqdm +from .. import filtering from .base import ( - SingleStepSampler, DPMPPStepMixin, HistorySingleStepSampler, ReversibleSingleStepSampler, + SingleStepSampler, registry, ) @@ -524,6 +525,147 @@ class RESMultistepStep(HistorySingleStepSampler, DPMPPStepMixin): yield from self.result(result, sigma_up, sigma_down=sigma_down) +# SEEDS-2 - Stochastic Explicit Exponential Derivative-free Solvers (VP Data Prediction) stage 2. +# arXiv: https://arxiv.org/abs/2305.14267 (NeurIPS 2023) +# Implementation referenced from ComfyUI. +class Seeds2Step(SingleStepSampler, DPMPPStepMixin): + name = "seeds_2" + self_noise = 3 + model_calls = 1 + allow_alt_cfgpp = False + uses_alt_noise = True + + def __init__(self, *args, r=0.5, **kwargs): + super().__init__(*args, **kwargs) + self.r = r + s2_options = self.options.get("seeds_2", {}) + sigma_blend_mode = s2_options.get("sigma_blend_mode", "lerp").strip() + self.sigma_blend_function = ( + filtering.BLENDING_MODES[sigma_blend_mode] + if sigma_blend_mode != "lerp" + else torch.lerp + ) + denoised_blend_mode = s2_options.get("denoised_blend_mode", "lerp").strip() + self.denoised_blend_function = ( + filtering.BLENDING_MODES[denoised_blend_mode] + if denoised_blend_mode != "lerp" + else torch.lerp + ) + self.disable_stage2_eta = bool(s2_options.get("disable_stage2_eta", False)) + stage2_stage1_noise_blend_mode = s2_options.get( + "stage2_stage1_noise_blend_mode", "lerp" + ).strip() + self.stage2_stage1_noise_blend_function = ( + filtering.BLENDING_MODES[stage2_stage1_noise_blend_mode] + if stage2_stage1_noise_blend_mode != "lerp" + else torch.lerp + ) + self.stage2_stage1_noise_ratio = s2_options.get( + "stage2_stage1_noise_ratio", 1.0 + ) + self.stage1_s_noise = s2_options.get("stage1_s_noise", 1.0) + self.stage2_s_noise = s2_options.get("stage2_s_noise", 1.0) + self.stage2_sigma_scale = s2_options.get("stage2_sigma_scale", 1.0) + + def step(self, x: torch.Tensor): + ss = self.ss + sigma = ss.sigma.to(dtype=torch.float64) + sigma_next = ss.sigma_next.to(dtype=torch.float64) + denoised = ss.denoised + + t_one = ss.sigma * 0 + 1.0 + + r, eta = self.r, self.get_dyn_eta() + fac = 1 / (2 * r) + lambda_s = ss.sigma_to_half_log_snr(sigma=sigma) + lambda_t = ss.sigma_to_half_log_snr(sigma=sigma_next) + h = lambda_t - lambda_s + h_eta = h * (eta + 1.0) + lambda_s_1 = self.sigma_blend_function( + lambda_s.unsqueeze(0), lambda_t.unsqueeze(0), r + ).squeeze(0) + sigma_s_1 = ss.half_log_snr_to_sigma(lambda_s_1) + + alpha_s_1 = sigma_s_1 * lambda_s_1.exp() + alpha_t = sigma_next * lambda_t.exp() + + s1_x_mult = sigma_s_1 / sigma * (-r * h * eta).exp() + s1_denoised_mult = alpha_s_1 * (-r * h_eta).expm1() + x_2 = ( + s1_x_mult.to(dtype=x.dtype) * x + - s1_denoised_mult.to(dtype=x.dtype) * denoised + ) + if eta != 0: + s1_noise_mult = (-2 * r * h * eta).expm1().neg().sqrt() + sde_noise1 = yield from self.result( + x_2 * 0, + s1_noise_mult.to(dtype=x.dtype), + sigma=ss.sigma, + sigma_next=sigma_s_1.to(dtype=x.dtype), + noise_sampler=self.alt_noise_sampler, + final=False, + ) + x_2 += sde_noise1 * (sigma_s_1.to(dtype=x.dtype) * self.stage1_s_noise) + + denoised_2 = self.call_model( + x_2, (sigma_s_1 * self.stage2_sigma_scale).to(dtype=x.dtype), call_index=1 + ).denoised + denoised_d = self.denoised_blend_function(denoised, denoised_2, fac) + + if self.disable_stage2_eta: + eta = 0.0 + h_eta = h + + s2_x_mult = sigma_next / sigma * (-h * eta).exp() + s2_denoised_mult = alpha_t * h_eta.neg().expm1() + x_curr = s2_x_mult.to(dtype=x.dtype) * x + x_curr -= s2_denoised_mult.to(dtype=x.dtype) * denoised_d + + if eta == 0: + return (yield from self.result(x_curr)) + + s2_s1_nr = self.stage2_stage1_noise_ratio + + segment_factor = ((r - 1.0) * h * eta).to(dtype=x.dtype) + s2_noise_mult = (segment_factor * 2.0).expm1().neg() ** 0.5 + sde_noise2_raw = yield from self.result( + x_curr * 0, + t_one, + sigma=sigma_s_1.to(dtype=x.dtype), + sigma_next=ss.sigma_next, + final=False, + ) + sde_noise2 = sde_noise2_raw * s2_noise_mult.to(dtype=x.dtype) + + if s2_s1_nr != 1.0: + # print( + # f"\n\nBLENDING: {s2_s1_nr:.4f}, {s1_noise_mult.item():.4f}, {s2_noise_mult.item():.4f}" + # ) + sde_noise1 = self.stage2_stage1_noise_blend_function( + ( + yield from self.result( + x_curr * 0, + t_one, + sigma=sigma_s_1.to(dtype=x.dtype), + sigma_next=ss.sigma_next, + final=False, + ) + ) + * s1_noise_mult.to(dtype=x.dtype), + sde_noise1, + s2_s1_nr, + ) + + sde_noise1 *= segment_factor.exp() + sde_noise2 += sde_noise1 + sde_noise2 *= ss.sigma_next * self.stage2_s_noise + x_curr += sde_noise2 + + yield from self.result( + x_curr, noise_scale=ss.sigma * 0, sigma_down=ss.sigma_next + ) + + registry.add( DEISStep, DPMPP2MSDEStep, @@ -538,4 +680,5 @@ registry.add( DPM2Step, DPMPP2SStep, RESMultistepStep, + Seeds2Step, ) diff --git a/py/step_samplers/misc.py b/py/step_samplers/misc.py index 04910e7..11c411b 100644 --- a/py/step_samplers/misc.py +++ b/py/step_samplers/misc.py @@ -10,7 +10,7 @@ class PingPongStep(SingleStepSampler): super().__init__(*args, **kwargs) pingpong_options = self.options.pop("pingpong", {}) self.pingpong_start_step = pingpong_options.get("start_step", 0) - self.pingpong_end_step = pingpong_options.get("end_step", 0) + self.pingpong_end_step = pingpong_options.get("end_step", 9999) def step(self, x): ss = self.ss diff --git a/py/substep_merging.py b/py/substep_merging.py index 47cc929..e31f97d 100644 --- a/py/substep_merging.py +++ b/py/substep_merging.py @@ -5,8 +5,7 @@ import tqdm from . import expression as expr from . import utils - -from .filtering import make_filter, FilterRefs, FILTER_HANDLERS +from .filtering import FILTER_HANDLERS, FilterRefs, make_filter from .noise import ImmiscibleNoise from .restart import Restart from .step_samplers import STEP_SAMPLERS @@ -451,6 +450,7 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler): s_noise=restart.get("s_noise", 1.0), custom_noise=restart_custom_noise, immiscible=restart.get("immiscible", False), + is_flow=ss.model.is_rectified_flow, ) def make_schedule(self, ss): @@ -505,10 +505,13 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler): if subss.idx >= max_idx: break if last_down is not None and last_down < ss.sigma_next: - restart_ns = self.restart.get_noise_sampler(ss.noise) - x += ss.noise.scale_noise( - restart_ns(last_down, ss.sigma_next, refs=ss.refs), - self.restart.get_noise_scale(last_down, ss.sigma_next), + x = self.restart.add_noise( + x, + sigma_from=last_down.item(), + sigma_to=ss.sigma_next.item(), + nsc=nsc, + refs=ss.refs, + in_place=True, ) pbar.update(0) return x @@ -653,11 +656,13 @@ class PingpongMergeSubstepsSampler(MergeSubstepsSampler): sigma_next, immiscible=fallback(self.immiscible, ss.noise.immiscible), ) - noise_refs = ss.refs | FilterRefs({ - "orig_x": orig_x, - "x": x, - "denoised": synth_denoised, - }) + noise_refs = ss.refs | FilterRefs( + { + "orig_x": orig_x, + "x": x, + "denoised": synth_denoised, + } + ) noise = ( noise_sampler(sigma, sigma_next, refs=noise_refs) * self.pingpong_s_noise ) diff --git a/py/substep_sampling.py b/py/substep_sampling.py index 37f64e1..f829bac 100644 --- a/py/substep_sampling.py +++ b/py/substep_sampling.py @@ -1,5 +1,6 @@ -import torch +from typing import NamedTuple +import torch from comfy.k_diffusion.sampling import get_ancestral_step from .filtering import FilterRefs @@ -7,6 +8,13 @@ from .model import History from .utils import fallback +class AncestralRatios(NamedTuple): + alpha_t: torch.Tensor + alpha_s: torch.Tensor + sigma_up: torch.Tensor + sigma_down: torch.Tensor + + class Items: def __init__(self, items=None): self.items = [] if items is None else items @@ -141,6 +149,10 @@ class SamplerState: self.substep = 0 self.total_steps = len(sigmas) - 1 self.cfg_scale_override = cfg_scale_override + self.is_flow = self.model.is_rectified_flow + self.offset_sigma = ( + model.model_sampling.percent_to_sigma(1e-04) if self.is_flow else None + ) self.update(idx) # Sets idx, sigma_prev, sigma, sigma_down, refs @property @@ -171,6 +183,23 @@ class SamplerState: def d(self): return self.hcur.d + # These two functions referenced from ComfyUI. + def sigma_to_half_log_snr( + self, *, sigma: torch.Tensor | None = None, idx: int | None = None + ) -> torch.Tensor: + if sigma is None and idx is None: + sigma = self.sigma + else: + sigma = sigma if sigma is not None else self.sigmas[idx] + if not self.is_flow: + return sigma.log().neg_() + if sigma.max() >= 1.0: + sigma = sigma * 0.0 + self.offset_sigma + return sigma.logit().neg_() + + def half_log_snr_to_sigma(self, half_log_snr: torch.Tensor) -> torch.Tensor: + return (torch.sigmoid if self.is_flow else torch.exp)(half_log_snr.neg()) + def update(self, idx=None, step=None, substep=None): idx = self.idx if idx is None else idx self.idx = idx @@ -185,6 +214,59 @@ class SamplerState: self.substep = substep self.refs = FilterRefs.from_ss(self) + def get_ancestral_step_ext( + self, + *, + sigma: torch.Tensor | None = None, + sigma_next: torch.Tensor | None = None, + eta: float = 1.0, + retry_increment: int = 0, + ): + sigma = fallback(sigma, self.sigma) + sigma_next = fallback(sigma_next, self.sigma_next) + sigma_empty = sigma_next * 0.0 + + def get_noeta_ratios(): + return AncestralRatios( + alpha_t=sigma_empty + 1.0, + alpha_s=sigma_empty + 1.0, + sigma_up=sigma_empty.clone(), + sigma_down=sigma_next.clone(), + ) + + if eta <= 0 or sigma_next.max().item() <= 1e-08: + return get_noeta_ratios() + orig_dtype = sigma.dtype + sigma = sigma.to(dtype=torch.float64) + sigma_next = sigma_next.to(dtype=torch.float64) + alpha_s = sigma * self.sigma_to_half_log_snr(sigma=sigma).exp() + alpha_t = sigma_next * self.sigma_to_half_log_snr(sigma=sigma_next).exp() + adj_sigma = sigma / alpha_s + adj_sigma_next = sigma_next / alpha_t + sd = su = None + while eta > 0: + sd, su = ( + v if isinstance(v, torch.Tensor) else sigma.new_full((1,), v) + for v in get_ancestral_step(adj_sigma, adj_sigma_next, eta=eta) + ) + if sd > 0 and su > 0: + break + else: + sd = su = None + if retry_increment <= 0: + break + # print(f"\nETA {eta} failed, retrying with {eta - retry_increment}") + eta -= retry_increment + if sd is None or su is None: + return get_noeta_ratios() + sd = alpha_t * sd + return AncestralRatios( + alpha_t=alpha_t.to(dtype=orig_dtype), + alpha_s=alpha_s.to(dtype=orig_dtype), + sigma_up=su.to(dtype=orig_dtype), + sigma_down=sd.to(dtype=orig_dtype), + ) + def get_ancestral_step( self, eta=1.0, sigma=None, sigma_next=None, retry_increment=0 ): @@ -265,13 +347,15 @@ class SamplerState: preview = (hi.x - hi.denoised) * 0.1 + hi.denoised else: preview = hi.denoised - return self.callback_({ - "x": hi.x, - "i": self.step, - "sigma": hi.sigma, - "sigma_hat": hi.sigma, - "denoised": preview, - }) + return self.callback_( + { + "x": hi.x, + "i": self.step, + "sigma": hi.sigma, + "sigma_hat": hi.sigma, + "denoised": preview, + } + ) def reset(self): self.hist.reset() diff --git a/py/triton_lsa.py b/py/triton_lsa.py new file mode 100644 index 0000000..64376be --- /dev/null +++ b/py/triton_lsa.py @@ -0,0 +1,386 @@ +import torch +import triton +import triton.language as tl + + +@triton.autotune( + configs=[ + triton.Config({"num_warps": 4, "num_stages": 2}, num_warps=4, num_stages=2), + triton.Config({"num_warps": 8, "num_stages": 2}, num_warps=8, num_stages=2), + triton.Config({"num_warps": 4, "num_stages": 3}, num_warps=4, num_stages=3), + triton.Config({"num_warps": 8, "num_stages": 3}, num_warps=8, num_stages=3), + ], + key=[ + "B", + "R", + "C", + "BLOCK_SIZE", + ], # Retune if matrix dimensions change significantly +) +@triton.jit +def auction_lap_kernel( + cost_ptr, + assign_ptr, + stride_b, + stride_r, + stride_c, + stride_assign_b, + stride_assign_r, + B: tl.constexpr, + R: tl.constexpr, + C: tl.constexpr, + epsilon, + max_iter, + BLOCK_SIZE: tl.constexpr, +): + pid = tl.program_id(0) + + cost_base = cost_ptr + pid * stride_b + assign_base = assign_ptr + pid * stride_assign_b + + offs = tl.arange(0, BLOCK_SIZE) + col_mask = offs < C + + # Prices and Owners in SRAM/Registers + prices = tl.zeros([BLOCK_SIZE], dtype=tl.float32) + owners = tl.full([BLOCK_SIZE], -1, dtype=tl.int32) + row_to_col = tl.full([BLOCK_SIZE], -1, dtype=tl.int32) + + iter_idx = 0 + unassigned_count = R + + loop_continue = tl.full([], 1, dtype=tl.int1) + + # Loop condition: + # 1. unassigned_count > 0: Logic handled inside, but we need a break mechanism + # 2. iter_idx < max_iter: Safety break + # 3. loop_continue: Did we make progress last time? + + while unassigned_count > 0 and iter_idx < max_iter and loop_continue: + # Reset progress flag + # loop_continue &= False + loop_continue = tl.full([], 0, dtype=tl.int1) + + # Gauss-Seidel pass over all rows + for i in tl.range(0, R): + # Check if row i is unassigned + curr_c = tl.sum(tl.where(offs == i, row_to_col, 0)) + + if curr_c == -1: + # Load costs + row_cost_ptr = cost_base + i * stride_r + offs + row_costs = tl.load(row_cost_ptr, mask=col_mask, other=-torch.inf) + + # Net Value + values = row_costs - prices + + # Find Best + best_val, best_idx = tl.max(values, axis=0, return_indices=True) + + # CRITICAL: Only proceed if this is a valid edge (not -inf) + if best_val > -torch.inf: + # We have a valid move, so we continue the outer loop + loop_continue = tl.full([], 1, dtype=tl.int1) + + # Find Second Best + mask_not_best = (offs != best_idx) & col_mask + vals_no_best = tl.where(mask_not_best, values, -torch.inf) + second_best_val = tl.max(vals_no_best, axis=0) + + # Compute Bid + bid = best_val - second_best_val + epsilon + + # Update Price + prices = tl.where(offs == best_idx, prices + bid, prices) + + # Update Owners + prev_owner = tl.sum(tl.where(offs == best_idx, owners, 0)) + + if prev_owner != -1: + # Kick out previous owner + row_to_col = tl.where(offs == prev_owner, -1, row_to_col) + unassigned_count += 1 + + # Assign to current row + owners = tl.where(offs == best_idx, i, owners) + row_to_col = tl.where(offs == i, best_idx, row_to_col) + unassigned_count -= 1 + + iter_idx += 1 + + # Store Result + store_offs = tl.arange(0, BLOCK_SIZE) + store_mask = store_offs < R + tl.store(assign_base + store_offs, row_to_col, mask=store_mask) + + +# ----------------------------------------------------------------------------- +# Python Helpers +# ----------------------------------------------------------------------------- + + +def rescale_simple( + t: torch.Tensor, + target_min: float = 0.0, + target_max: float = 1.0, + *, + start_dim: int = 1, + eps: float = 1e-07, +) -> torch.Tensor: + width = target_max - target_min + if width == 0.0: + return torch.zeros_like(t) + orig_shape = t.shape + t = t.flatten(start_dim=start_dim) + min_val, max_val = t.aminmax(dim=-1, keepdim=True) + normalized = t - min_val + normalized /= (max_val - min_val).add_(eps) + normalized *= width + if target_min != 0.0: + normalized += target_min + return normalized.clamp_(target_min, target_max).reshape(orig_shape) + + +def _greedy_fill_missing(assignments: torch.Tensor, C: int) -> None: + """ + Fills unassigned rows (-1) in the assignments tensor with available columns. + This acts as a fallback when the Auction algorithm hits max_iter without + full convergence. + + Args: + assignments: Tensor of shape (B, R) containing col indices or -1. + C: Total number of columns available. + """ + # Identify which batch items have unassigned rows + # This is usually a very small subset (e.g., < 1% of the batch) + problem_batches = (assignments == -1).any(dim=1).nonzero().flatten() + + if problem_batches.numel() == 0: + return + + device = assignments.device + + # Iterate only over the problematic batch items + # (Looping is acceptable here as B_subset is typically tiny) + for b_idx in problem_batches: + # 1. Find which rows are missing an assignment + row_mask = assignments[b_idx] == -1 + missing_rows = row_mask.nonzero().flatten() + n_needed = missing_rows.shape[0] + + # 2. Find which columns are already used + used_cols = assignments[b_idx][~row_mask] + + # 3. Find free columns (Set difference: All - Used) + # Create a boolean mask of all columns, then mark used ones as False + # efficient on GPU for mid-sized C + col_mask = torch.ones(C, device=device, dtype=torch.bool) + col_mask[used_cols.long()] = False + + free_cols = col_mask.nonzero().flatten() + + # 4. Assign the first N free columns to the N missing rows + # Since R <= C in this context (due to transpose logic in wrapper), + # free_cols.numel() is guaranteed to be >= n_needed. + assignments[b_idx, missing_rows] = free_cols[:n_needed].to(assignments.dtype) + + +def batch_linear_assignment( + cost_matrix: torch.Tensor, + *, + maximize: bool = False, + max_iter: int | None = None, + fill_missing: bool = True, + rescale_costs: tuple[float, float] | None = (0.0, 1.0), + invert_costs_mode: bool = True, + eps: float = 1e-3, +): + if cost_matrix.ndim != 3: + raise ValueError("Cost matrix must be (B, R, C)") + if not cost_matrix.is_cuda: + raise ValueError("Cost matrix must be a CUDA tensor") + if not cost_matrix.is_contiguous(): + raise ValueError("Cost matrix must be contiguous") + + B, R, C = cost_matrix.shape + device = cost_matrix.device + + # 1. Handle Rectangular Matrices + # The Auction algorithm assigns Rows -> Cols. + # It naturally handles R <= C (finding best col for every row). + # If R > C, we must transpose to match Cols -> Rows, then invert the result. + if R > C: + transposed = True + cost_matrix = cost_matrix.mT.contiguous() + # Swap R and C for the kernel execution + R, C = C, R + else: + transposed = False + + if rescale_costs is not None: + cost_matrix = rescale_simple(cost_matrix, *rescale_costs) + # Note: We use float32 for atomic compatibility and speed + cost_matrix = cost_matrix.to(torch.float32, copy=rescale_costs is None) + + if not maximize: + # Maximize (Value - Price) -> Minimize Cost + cost_matrix = cost_matrix.neg_() + if invert_costs_mode and rescale_costs is not None: + cost_matrix += sum(rescale_costs) + + assignments = torch.full( + (B, R), + -1, + device=device, + dtype=torch.int32, + ) + + max_dim = max(R, C) + BLOCK_SIZE = max(32, triton.next_power_of_2(max_dim)) + + # Safety limit + max_iter = max_iter if max_iter is not None else int(max(2000, R * C)) + + grid = (B,) + + auction_lap_kernel[grid]( + cost_matrix, + assignments, + cost_matrix.stride(0), + cost_matrix.stride(1), + cost_matrix.stride(2), + assignments.stride(0), + assignments.stride(1), + B, + R, + C, + eps, + max_iter, + BLOCK_SIZE=BLOCK_SIZE, + ) + + if fill_missing: + _greedy_fill_missing(assignments, C) + + assignments = assignments.long() + + if not transposed: + return assignments + + # 2. Post-process Rectangular Results + # We computed Col -> Row. We need Row -> Col. + # assignments shape is currently (B, Original_Cols) + # We want output shape (B, Original_Rows) + + real_rows = C # C is the 'large' dimension (Original Rows) + output = torch.full((B, real_rows), -1, device=device, dtype=torch.long) + + # Create indices for the scatter source + # We want: output[row_idx] = col_idx + # Currently we have: assignments[col_idx] = row_idx + src_col_indices = torch.arange(R, device=device).unsqueeze(0).expand(B, R) + + # We use scatter. index=assignments (the rows), src=col_indices + # To handle -1s in assignments, we clamp to 0 and then mask the result + safe_assigns = assignments.clamp(min=0) + output.scatter_(1, safe_assigns, src_col_indices) + + # Cleanup: Any row that wasn't targeted by the scatter should be -1 + # The scatter might have written to index 0 if assignment was -1 + # Re-verify logic: + for b in range(B): + valid_mask = assignments[b] >= 0 + # Reset output + output[b].fill_(-1) + # Only write valid mappings + # output[b, row_id] = col_id + output[b, assignments[b, valid_mask]] = src_col_indices[b, valid_mask] + + return output + + +def assignments_to_indices( + assignments: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """ + Converts a dense assignment tensor (from Triton/Hungraian) to + batched SciPy-style indices. + + Args: + assignments (torch.Tensor): Shape (B, R). Values are col indices or -1. + + Returns: + row_ind (torch.Tensor): Shape (B, K) where K = min(R, C). + col_ind (torch.Tensor): Shape (B, K). + """ + B, R = assignments.shape + device = assignments.device + + # 1. Create a mask of valid assignments (values >= 0) + # In a rectangular assignment, the number of valid matches + # is always min(Rows, Cols). + mask = assignments >= 0 + + # 2. Extract Column Indices + # We select the values from the assignment tensor that are valid. + # We reshape to (B, -1) to preserve the batch dimension. + col_ind = assignments[mask].view(B, -1) + + # 3. Extract Row Indices + # We need a grid of row indices [0, 1, 2, ... R-1] repeated B times + row_grid = ( + torch.arange(R, device=device, dtype=assignments.dtype) + .unsqueeze(0) + .expand(B, R) + ) + row_ind = row_grid[mask].view(B, -1) + + return row_ind, col_ind + + +def batch_linear_assignment_shuffled( + cost_matrix: torch.Tensor, + *args, + **kwargs: dict, +) -> torch.Tensor: + generator = kwargs.pop("generator", None) + # cost_matrix: [B, R, C] + B, R = cost_matrix.shape[:2] + + # 1. Generate a random permutation for the rows + # We use one perm for the whole batch for efficiency, + # or you can do it per-batch-item if B is small and quality is critical. + # Here we shuffle all rows commonly. + perm = torch.randperm( + R, + device=cost_matrix.device, + generator=generator, + ) + + # 2. Shuffle the input (Row dimension is dim 1) + # This creates a shuffled view/copy of the cost matrix + shuffled_cost = cost_matrix[:, perm, :] + + # 3. Run the Solver + shuffled_assignments = batch_linear_assignment( + shuffled_cost, + *args, + **kwargs, + ) # Returns [B, R] + + # 4. Un-shuffle the results + # We need to map the results back to their original row positions. + # shuffled_assignments[b, i] corresponds to the row 'perm[i]' + # We want final_assignments[b, perm[i]] = shuffled_assignments[b, i] + + # Create the inverse permutation or just scatter back + final_assignments = torch.empty_like(shuffled_assignments) + + # Expand perm for the batch: [B, R] + batch_perm = perm.unsqueeze(0).expand(B, R) + + # Scatter the results back to original positions + # dim=1, index=batch_perm, src=shuffled_assignments + final_assignments.scatter_(1, batch_perm, shuffled_assignments) + + return final_assignments diff --git a/py/utils.py b/py/utils.py index 0dd2c2c..bc46219 100644 --- a/py/utils.py +++ b/py/utils.py @@ -1,7 +1,10 @@ +from __future__ import annotations + import contextlib +import math +from functools import partial import torch - from comfy.k_diffusion.sampling import to_d # def scale_noise_( @@ -39,25 +42,118 @@ from comfy.k_diffusion.sampling import to_d def scale_noise( - noise, - factor=1.0, + noise: torch.Tensor, + factor: float = 1.0, *, - normalized=True, - normalize_dims=(-3, -2, -1), -): + normalized: bool = True, + normalize_dims: tuple[int, ...] = (-3, -2, -1), + eps: float = 1e-08, +) -> torch.Tensor: if not normalized or noise.numel() == 0: return noise * factor if factor != 1 else noise - noise = noise / noise.std(dim=normalize_dims, keepdim=True) - return noise.sub_(noise.mean(dim=normalize_dims, keepdim=True)).mul_(factor) + std = noise.std(dim=normalize_dims, keepdim=True) + noise = noise / torch.where(std != 0.0, std, eps) + noise -= noise.mean(dim=normalize_dims, keepdim=True) + return noise if factor == 1.0 else noise.mul_(factor) + + +def range_wrap( + x: torch.Tensor, + min_val: float | torch.Tensor, + max_val: float | torch.Tensor, +) -> torch.Tensor: + return min_val + (x - min_val).remainder_(max_val - min_val) def _quantile_norm_scaledown( noise: torch.Tensor, nq: torch.Tensor, + *, + dim, **_kwargs: dict, ) -> torch.Tensor: - mv = noise.abs().max().detach().item() - return noise if mv == 0 else torch.where(noise.abs() > nq, noise * (nq / mv), noise) + noiseabs = noise.abs() + mv = noiseabs.max(dim=dim, keepdim=True).values.clamp(min=1e-06) + return ( + noise + if mv.sum().item() == 0 + else torch.where(noiseabs > nq, noise * (nq / mv), noise) + ) + + +def _quantile_norm_wave( + noise: torch.Tensor, + nq: torch.Tensor, + *, + preserve_sign: bool = False, + wave_function=torch.sin, + pi_factor: float = 0.5, + wrong_mode: bool = False, + **_kwargs: dict, +) -> torch.Tensor: + if wrong_mode: + multiplier = 1.0 / ((math.pi * pi_factor) / nq) + else: + multiplier = 1.0 / (nq / (math.pi * pi_factor)) + pos_mask = noise >= 0 + neg_mask = ~pos_mask + result = torch.zeros_like(noise) + result[pos_mask] = wave_function(noise.mul(multiplier))[pos_mask] + result[neg_mask] = wave_function(noise.mul(multiplier))[neg_mask] + result *= nq + return result.copysign(noise) if preserve_sign else result + + +def _quantile_norm_mode( + noise: torch.Tensor, + nq: torch.Tensor, + *, + dim: int | None, + decimals=1, + **_kwargs: dict, +) -> torch.Tensor: + return torch.where( + noise.abs() > nq, + noise.round(decimals=decimals).mode(dim=dim, keepdim=True).values, + noise, + ) + + +def _quantile_norm_replace( + noise: torch.Tensor, + nq: torch.Tensor, + *, + keep_sign: bool = False, + avoid_sign: bool = False, + count: int = 1, + count_flipping: bool = False, + **_kwargs: dict, +) -> torch.Tensor: + mask = noise.abs() <= nq + candidates = noise[mask].flatten() + n_candidates = candidates.numel() + idxs = torch.arange(noise.numel()) % n_candidates + cresult = candidates[idxs] + if count < 2: + candidates = cresult + else: + multiplier = 1.0 / count + cresult = cresult * multiplier # noqa: PLR6104 + for i in range(1, count): + cresult += ( + candidates[ + torch.roll( + idxs, + i if not count_flipping or (i % 2) == 0 else -i, + dims=(-1,), + ) + ] + * multiplier + ) + candidates = cresult.reshape(noise.shape) + if keep_sign or avoid_sign: + candidates = candidates.copysign_(noise.neg() if avoid_sign else noise) + return torch.where(mask, noise, candidates) quantile_handlers = { @@ -69,14 +165,66 @@ quantile_handlers = { noise.tanh().mul_(nq.abs()), noise, ), - "sigmoid": lambda noise, nq, **_kwargs: noise.sigmoid() - .mul_(nq.abs()) - .copysign(noise), + "sigmoid_keepsign": lambda noise, nq, **_kwargs: ( + noise.sigmoid().mul_(nq.abs()).copysign(noise) + ), + "sigmoid": lambda noise, nq, **_kwargs: ( + noise.sigmoid().mul_(nq.abs() * 2).sub_(nq.abs()) + ), "sigmoid_outliers": lambda noise, nq, **_kwargs: torch.where( noise.abs() > nq, noise.sigmoid().mul_(nq.abs()).copysign(noise), noise, ), + "sin": partial(_quantile_norm_wave, wave_function=torch.sin), + "sin_wholepi": partial( + _quantile_norm_wave, + wave_function=torch.sin, + pi_factor=1.0, + ), + "sin_keepsign": partial( + _quantile_norm_wave, + wave_function=torch.sin, + preserve_sign=True, + ), + "sin_wrong": partial(_quantile_norm_wave, wave_function=torch.sin, wrong_mode=True), + "sin_wrong_wholepi": partial( + _quantile_norm_wave, + wave_function=torch.sin, + pi_factor=1.0, + wrong_mode=True, + ), + "sin_wrong_keepsign": partial( + _quantile_norm_wave, + wave_function=torch.sin, + preserve_sign=True, + wrong_mode=True, + ), + "cos": partial(_quantile_norm_wave, wave_function=torch.cos), + "cos_wholepi": partial( + _quantile_norm_wave, + wave_function=torch.cos, + pi_factor=1.0, + ), + "cos_keepsign": partial( + _quantile_norm_wave, + wave_function=torch.cos, + preserve_sign=True, + ), + "cos_wrong": partial(_quantile_norm_wave, wave_function=torch.cos, wrong_mode=True), + "cos_wrong_wholepi": partial( + _quantile_norm_wave, + wave_function=torch.cos, + pi_factor=1.0, + wrong_mode=True, + ), + "cos_wrong_keepsign": partial( + _quantile_norm_wave, + wave_function=torch.cos, + preserve_sign=True, + wrong_mode=True, + ), + "atan": lambda noise, nq, **_kwargs: noise.atan().mul_(nq.abs() / (math.pi / 2)), "tenth": lambda noise, nq, **_kwargs: torch.where( noise.abs() > nq, noise * 0.1, @@ -93,6 +241,80 @@ quantile_handlers = { noise, 0, ), + "mean": lambda noise, nq, *, dim, **_kwargs: torch.where( + noise.abs() > nq, + noise.mean(dim=dim, keepdim=True), + noise, + ), + "median": lambda noise, nq, *, dim, **_kwargs: torch.where( + noise.abs() > nq, + noise.median(dim=dim, keepdim=True).values, + noise, + ), + "mode_1dec": partial(_quantile_norm_mode, decimals=1), + "mode_2dec": partial(_quantile_norm_mode, decimals=2), + "replace": _quantile_norm_replace, + "replace_keepsign": partial(_quantile_norm_replace, keep_sign=True), + "replace_avoidsign": partial(_quantile_norm_replace, avoid_sign=True), + "replace_2pt": partial(_quantile_norm_replace, count=2), + "replace_3pt": partial(_quantile_norm_replace, count=3), + "replace_2pt_flip": partial(_quantile_norm_replace, count=2, count_flipping=True), + "replace_3pt_flip": partial(_quantile_norm_replace, count=3, count_flipping=True), + "replace_2pt_keepsign": partial( + _quantile_norm_replace, + count=2, + keep_sign=True, + ), + "replace_3pt_keepsign": partial( + _quantile_norm_replace, + count=3, + keep_sign=True, + ), + "replace_2pt_flip_keepsign": partial( + _quantile_norm_replace, + count=2, + count_flipping=True, + keep_sign=True, + ), + "replace_3pt_flip_keepsign": partial( + _quantile_norm_replace, + count=3, + count_flipping=True, + keep_sign=True, + ), + "replace_2pt_avoidsign": partial( + _quantile_norm_replace, + count=2, + avoid_sign=True, + ), + "replace_3pt_avoidsign": partial( + _quantile_norm_replace, + count=3, + avoid_sign=True, + ), + "replace_2pt_flip_avoidsign": partial( + _quantile_norm_replace, + count=2, + count_flipping=True, + avoid_sign=True, + ), + "replace_3pt_flip_avoidsign": partial( + _quantile_norm_replace, + count=3, + count_flipping=True, + avoid_sign=True, + ), + "wrap": lambda noise, nq, **_kwargs: range_wrap(noise, -nq, nq), + "wrap_keepsign": lambda noise, nq, **_kwargs: torch.where( + noise.abs() > nq, + range_wrap(noise, -nq, nq).copysign_(noise), + noise, + ), + "wrap_avoidsign": lambda noise, nq, **_kwargs: torch.where( + noise.abs() > nq, + range_wrap(noise, -nq, nq).copysign_(noise.neg()), + noise, + ), } @@ -100,48 +322,40 @@ quantile_handlers = { def quantile_normalize( noise: torch.Tensor, *, - quantile: float = 0.75, + quantile: float | tuple | list = 0.75, dim: int | None = 1, flatten: bool = True, nq_fac: float = 1.0, pow_fac: float = 0.5, strategy: str = "clamp", strategy_handler=None, + eps=1e-08, ) -> torch.Tensor: - if quantile is None or quantile <= 0 or quantile >= 1: + if noise.numel() == 0: return noise - orig_shape = noise.shape if isinstance(quantile, (tuple, list)): - quantile = torch.tensor( - quantile, - device=noise.device, - dtype=noise.dtype, - ) - qdim = dim - if noise.ndim > 1 and flatten: - if qdim is not None and qdim >= noise.ndim: - qdim = 1 if noise.ndim > 2 else None - if qdim is None: - flatdim = 0 - elif -1 < qdim < 2: # 0, 1 - flatdim = qdim + 1 - elif 1 < qdim < 4: # 2, 3 - noise = noise.movedim(qdim, 1) - tempshape = noise.shape - flatdim = 2 - else: - raise ValueError( - "Cannot handling quantile normalization flattening dims > 3", + for q in quantile: + noise = quantile_normalize( + noise=noise, + quantile=q, + dim=dim, + flatten=flatten, + nq_fac=nq_fac, + pow_fac=pow_fac, + strategy=strategy, + strategy_handler=strategy_handler, ) + return noise + if quantile is None or quantile >= 1 or quantile <= -1: + return noise + centered = quantile < 0 + absquantile = abs(quantile) + orig_shape = noise.shape + if noise.ndim > 1 and flatten: + flatnoise = noise.flatten(start_dim=dim) else: - flatdim = None - nq = torch.quantile( - (noise if flatdim is None else noise.flatten(start_dim=flatdim)).abs(), - quantile, - dim=-1, - ) - nq_shape = tuple(nq.shape) + (1,) * (noise.ndim - nq.ndim) - nq = nq.mul_(nq_fac).reshape(*nq_shape) + flatten = False + flatnoise = noise handler = ( quantile_handlers.get(strategy) if strategy_handler is None @@ -149,18 +363,45 @@ def quantile_normalize( ) if handler is None: raise ValueError("Unknown strategy") - noise = handler( - noise, - nq, - dim=dim, - flatten=flatten, - ) - noise = noise.abs().pow(pow_fac).copysign(noise) - if flatdim is not None and qdim in {2, 3}: - return ( - noise.reshape(tempshape).movedim(1, qdim).reshape(orig_shape).contiguous() + if not centered: + nq = torch.quantile( + flatnoise.abs(), + quantile, + dim=-1 if flatten else dim, + keepdim=True, ) - return noise + nq = nq.mul_(nq_fac).add_(eps) + # print(f"\nNQ: {nq}") + noise = handler( + flatnoise, + nq, + orig_noise=noise, + dim=dim, + flatten=flatten, + ) + else: + absnoise = flatnoise.abs() + maxabs = absnoise.amax(dim=-1 if flatten else dim, keepdim=True) + proxy = flatnoise.sign().mul_(maxabs - absnoise) + nq_proxy = torch.quantile( + proxy.abs(), + absquantile, + dim=-1 if flatten else dim, + keepdim=True, + ) + nq_proxy = nq_proxy.mul_(nq_fac).add_(eps) + # print(f"\nNQ proxy: {nq_proxy}") + out_proxy = handler( + proxy, + nq_proxy, + orig_noise=noise, + dim=dim, + flatten=flatten, + ) + noise = out_proxy.sign().mul_(maxabs - out_proxy.abs()) + if pow_fac not in {0.0, 1.0}: + noise = noise.abs().pow_(pow_fac).copysign(noise) + return noise if noise.shape == orig_shape else noise.reshape(orig_shape) # def scale_noise( diff --git a/py/wavelet_functions.py b/py/wavelet_functions.py new file mode 100644 index 0000000..638f2a0 --- /dev/null +++ b/py/wavelet_functions.py @@ -0,0 +1,238 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Callable + +import torch + +from .utils import fallback + +if TYPE_CHECKING: + from collections.abc import Sequence + +try: + import pytorch_wavelets as ptwav + import pywt + + HAVE_WAVELETS = True +except ImportError: + ptwav = None + pywt = None + HAVE_WAVELETS = False + + +class Wavelet: + DEFAULT_MODE = "symmetric" + DEFAULT_LEVEL = 3 + DEFAULT_WAVE = "db4" + DEFAULT_USE_1D_DWT = False + DEFAULT_USE_DTCWT = False + DEFAULT_QSHIFT = "qshift_a" + DEFAULT_BIORT = "near_sym_a" + + def __init__( + self, + *, + wave: str = DEFAULT_WAVE, + level: int = DEFAULT_LEVEL, + mode: str = DEFAULT_MODE, + use_1d_dwt: bool = DEFAULT_USE_1D_DWT, + use_dtcwt: bool = DEFAULT_USE_DTCWT, + biort: str = DEFAULT_BIORT, + qshift: str = DEFAULT_QSHIFT, + inv_wave: str | None = None, + inv_mode: str | None = None, + inv_biort: str | None = None, + inv_qshift=None, + device: str | torch.device | None = None, + ): + if not HAVE_WAVELETS: + raise RuntimeError( + "Wavelet noise requires the pytorch_wavelets package to be installed in your Python environment", + ) + inv_wave = fallback(inv_wave, wave) + inv_mode = fallback(inv_mode, mode) + inv_biort = fallback(inv_biort, biort) + inv_qshift = fallback(inv_qshift, qshift) + if use_dtcwt: + fwdfun, invfun = ptwav.DTCWTForward, ptwav.DTCWTInverse + elif use_1d_dwt: + fwdfun, invfun = ptwav.DWT1DForward, ptwav.DWT1DInverse + else: + fwdfun, invfun = ptwav.DWTForward, ptwav.DWTInverse + if use_dtcwt: + self._wavelet_forward = fwdfun( + J=level, + mode=mode, + biort=biort, + qshift=qshift, + ) + self._wavelet_inverse = invfun( + mode=inv_mode, + biort=inv_biort, + qshift=inv_qshift, + ) + else: + self._wavelet_forward = fwdfun(J=level, wave=wave, mode=mode) + self._wavelet_inverse = invfun(wave=inv_wave, mode=inv_mode) + if device is not None: + self._wavelet_forward = self._wavelet_forward.to(device=device) + self._wavelet_inverse = self._wavelet_inverse.to(device=device) + + def forward( + self, + t: torch.Tensor, + *, + forward_function: Callable | None = None, + ) -> tuple[torch.Tensor, tuple]: + return fallback(forward_function, self._wavelet_forward)(t) + + def inverse( + self, + yl: torch.Tensor, + yh: tuple, + *, + inverse_function: Callable | None = None, + two_step_inverse: bool = False, + ) -> torch.Tensor: + inverse_function = fallback(inverse_function, self._wavelet_inverse) + if not two_step_inverse: + return inverse_function((yl, yh)) + result = inverse_function((torch.zeros_like(yl), yh)) + result += inverse_function(( + yl, + tuple(torch.zeros_like(yh_band) for yh_band in yh), + )) + return result + + def to(self, *args: list, copy: bool = False, **kwargs: dict) -> Wavelet: + o = Wavelet.__new__(Wavelet) if copy else self + o._wavelet_forward = self._wavelet_forward.to(*args, **kwargs) # noqa: SLF001 + o._wavelet_inverse = self._wavelet_inverse.to(*args, **kwargs) # noqa: SLF001 + return o + + @staticmethod + def wavelist() -> tuple: + return tuple(pywt.wavelist()) if HAVE_WAVELETS else () + + @staticmethod + def biortlist() -> tuple: + return ( + ("near_sym_a", "near_sym_b", "antonini", "legall") if HAVE_WAVELETS else () + ) + + @staticmethod + def qshiftlist() -> tuple: + return ( + ("qshift_a", "qshift_b", "qshift_c", "qshift_d", "qshift_06") + if HAVE_WAVELETS + else () + ) + + @staticmethod + def modelist() -> tuple: + return ( + ( + "symmetric", + "zero", + "reflect", + "replicate", + "periodization", + "periodic", + "constant", + ) + if HAVE_WAVELETS + else () + ) + + +def expand_yh_scales( + yh: Sequence, + *, + yh_scales: float | Sequence = 1.0, +) -> float | tuple: + yhlen = len(yh) + yh_shape = yh[0].shape + # Doesn't make sense to target orientations for 1D DWD (3D here). + olen = yh_shape[2] if len(yh_shape) > 3 else 1 + # print(f"\nSIZES: yhlen={yhlen}, olen={olen}, yh_shape={yh[0].shape}") + if isinstance(yh_scales, (float, int)): + return ((float(yh_scales),) * olen,) * yhlen + otemplate = (1.0,) * olen + yh_scales = tuple( + (float(band),) * olen + if isinstance(band, (float, int)) + else ( + ( + *(float(i) for i in band[:olen]), + *otemplate[: olen - len(band[:olen])], + ) + if isinstance(band, (tuple, list)) + else band + ) + for band in yh_scales + ) + if "fill" in yh_scales: + fillidx = yh_scales.index("fill") + if "fill" in yh_scales[fillidx + 1 :]: + raise ValueError("Only one fill allowed.") + if fillidx == 0 or len(yh_scales) < 2: + raise ValueError( + "Invalid fill value, cannot be in the first position or the only item.", + ) + yhslen = len(yh_scales) + if yhslen - 1 < yhlen: + # Need to pad. + fill = (yh_scales[fillidx - 1],) * (yhlen - (len(yh_scales) - 1)) + yh_scales = (*yh_scales[:fillidx], *fill, *yh_scales[fillidx + 1 :]) + else: + # Just remove the "fill". + yh_scales = (*yh_scales[:fillidx], *yh_scales[fillidx + 1 :]) + return yh_scales[:yhlen] + + +def wavelet_scaling( + yl: torch.Tensor, + yh: Sequence, + yl_scale: float | torch.Tensor, + yh_scales: float | Sequence | None, + *, + in_place: bool = False, +) -> tuple: + if not in_place: + yl = yl.clone() + yh = tuple(yhband.clone() for yhband in yh) + if yl_scale != 1.0: + yl *= yl_scale + yh_scales = expand_yh_scales( + yh, + yh_scales=yh_scales if yh_scales is not None else 1.0, + ) + for hscale, ht in zip(yh_scales, yh): + if isinstance(hscale, (int, float)): + ht *= hscale # noqa: PLW2901 + continue + for lidx in range(min(ht.shape[2], len(hscale))): + ht[:, :, lidx] *= hscale[lidx] + return (yl, yh) + + +def wavelet_blend( + a: tuple, + b: tuple, + *, + yl_factor: torch.Tensor | float, + blend_function: Callable, + yh_factor: torch.Tensor | float | None = None, + yh_blend_function: Callable | None = None, +) -> tuple: + if not isinstance(yl_factor, torch.Tensor): + yl_factor = a[0].new_full((1,), yl_factor) + if yh_factor is None: + yh_factor = yl_factor + elif not isinstance(yh_factor, torch.Tensor): + yh_factor = a[0].new_full((1,), yh_factor) + yh_blend_function = fallback(yh_blend_function, blend_function) + return ( + blend_function(a[0], b[0], yl_factor), + tuple(yh_blend_function(ta, tb, yh_factor) for ta, tb in zip(a[1], b[1])), + )