Author SHA1 Message Date
blepping a4fed311a8 Partial 5D latent support for custom noise types 2025-02-27 12:33:35 -07:00
blepping 2988afa34a Fix Bleh and Restart integration 2025-02-17 03:09:46 -07:00
blepping 68fc7418d1 Yet another noise generation refactor (#12)
* Refactor noise generation
Try to make option passing and CPU/GPU noise selection work
Add advanced custom noise node that allows for parameter passing
Add wavelet noise type

* Add WaveletFilteredNoise node, other fixes

* Fix Brownian arg passing

* Generalized distribution noise for most torch.distributions

* Distro noise improvements, add SonarAdvancedDistroNoise node

* More distributions!

* Add SonarResizedNoise node

* Momentum sampler refactor/improvements (I hope)

* Better approach to integration with external nodes
Documentation updates
Other cleanups

* Internal cleanups and refactoring.
Some integration improvements.
Bump date in changelog

* Add round and step to node FLOAT inputs that did not have it
2025-01-31 17:33:04 -07:00
blepping 6d15c0bbca Use ComfyUI union types for wildcard inputs when available 2024-12-05 03:57:16 -07:00
blepping f7cbbfcbda Merge pull request #11 from blepping/nov2024update
November 2024 mega update
2024-11-30 01:01:33 -07:00
12 changed files with 3432 additions and 1258 deletions
+43
View File
@@ -64,6 +64,47 @@ You can optionally plug this into the Sonar sampler nodes. See the [Guidance](#g
Very abbreviated section. The init type can make a big difference. If you use `RANDOM` you can get away with setting `direction` to high values (like up to `2.25` or so) and absurdly low values (like `-30.0`). It's also possible to set `momentum` and `momentum_hist` to negative values, although whether it's a good idea...
<details>
<summary>Click to expand advanced parameters info</summary>
There are some extra advanced parameters that may be passed by YAML/JSON using `SamplerConfigOVerride`'s `yaml_parameters`. Defaults:
```yaml
sonar_params:
# One of: classic, new, denoised
# classic: Should be the same as the way it works in the A1111 extension.
# new: Possibly improved version that doesn't blend in the history again.
# denoised: Instead of using the noise prediction, we do momentum on denoised instead.
momentum_mode: new
# The following two parameters may be used to control when
# momentum sampling is active. Steps are 0-based with 0 being the first step.
momentum_start_step: 0
momentum_end_step: 9999
# Controls whether history always gets updated, whether or not within the
# start/end step range or only in that range. Can be used to affect the initial
# history value.
always_update_history: true
# Only applies when the init type is RAND.
rand_init_noise_multiplier: 1.0
# If you have ComfyUI-bleh installed, you can use any blend mode it provides.
# Otherwise you can have your blend mode in any color you want as long as it's lerp.
blend_mode: lerp
# Defaultss to blend_mode if unset.
momentum_blend_mode: null
# Defaults to blend_mode if unset. Only applies to linear guidance mode.
guidance_blend_mode: null
```
Additionally, it's possible to override the normal Sonar parameters here as well. If they exist in the `sonar_params` block, they will overwrite the values in the node.
</details>
## Guidance
You can try the `SamplerSonarNaive` sampler which has an optional latent input. The guidance _probably_ isn't working correctly and the implementation definitely isn't exactly the same as the original A1111 version but it still might be fun to play with. The `linear` guidance type is a lot more sensitive to the `guidance_factor` than the `euler` type. For `euler`, reasonable values are around `0.01` to `0.1`, for `linear` reasonable values are more like `0.001` to `0.02`. It is also possible to set guidance factor to a negative value, I've found this results in high contrast and very vivid colors.
@@ -108,9 +149,11 @@ Original Sonar Sampler implementation (for A1111): https://github.com/Kahsolt/st
My version was initially based on this Sonar sampler implementation for Diffusers: https://github.com/alexblattner/modified-euler-samplers-for-sonar-diffusers/
* Many noise generation functions copied from https://github.com/Clybius/ComfyUI-Extra-Samplers with only minor modifications. I may have broken some of them in the process _or_ they may not have been suitable for use and I took them anyway. If they don't work it is not a reflection on the original source.
* Noise spectral modulation modified from https://github.com/Clybius/ComfyUI-Extra-Samplers
* New pyramid noise based on implementation in [Jonathan Whitaker](https://wandb.ai/johnowhitaker/multires_noise/reports/Multi-Resolution-Noise-for-Diffusion-Model-Training--VmlldzozNjYyOTU2)'s article on multi-resolution noise.
* Original `SonarPowerNoise` contributed by [elias-gaeros](https://github.com/elias-gaeros/). Additionally, he provided a lot of guidance with refactoring it to allow separate filtering and other enhancements and answered a multitude of dumb questions. To say those changes are only co-authored is probably giving myself too much credit. Thank you! Your patience and help is very much appreciated.
* New 1/f (onef) and power law (white, grey, violet, velvet) noise types referenced from https://github.com/WASasquatch/PowerNoiseSuite
* Wavelet noise idea (and some of the default settings) from https://github.com/ClownsharkBatwing/RES4LYF
## Errata
+20
View File
@@ -2,6 +2,26 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 20250227
* Add 5D latent (video models) support for most custom noise types.
## 20250130
*Note*: May change seeds.
This set of changes includes some pretty major internal refactoring. Definitely possible that I broke something, so please create an issue if you run into problems.
* Noise generation should now respect whether generating on CPU vs GPU is selected. Previously it likely was defaulting to generating on GPU. This may change seeds.
* Refactored momentum samplers, this may change seeds especially if you were using weird parameters like negative direction.
* Added some new parameters for momentum samplers.
* Removed the `s_noise` and churn parameters from the normal Sonar Euler sampler. May break workflows. (Churn was the predecessor to ancestral samplers and is basically obsolete.)
* Added `wavelet` and `distro` noise types.
* Added `SonarCustomNoiseAdv` node that allows passing parameters via YAML.
* Added `SonarResizedNoise` node that allows you to generate noise at a fixed size and then crop/resize it to match the generation.
* Added `SonarAdvancedDistroNoise` node that allows generating noise with basically all the distributions PyTorch supports.
* Added `SonarWaveletFilteredNoise` node that lets you filter another noise generator using wavelets.
## 20241129
*Note*: Contains some potentially workflow-breaking changes.
+68
View File
@@ -52,6 +52,21 @@ Parameters:
***
### `SonarCustomNoiseAdv`
Same as the `SonarCustomNoise` except it also includes a text widget for passing parameters by YAML or JSON (JSON is valid YAML).
Just for example, instead of using the absurdly large `SonarAdvancedDistroNoise` node, you could do something like:
```yaml
distro: wishart
quantile_norm: 0.5
wishart_cov_size: 4
wishart_df: 3.5
```
***
### `NoisyLatentLike`
This node takes a reference latent and generates noise of the same shape. The one required input is `latent`.
@@ -108,6 +123,59 @@ More extensive documentation TBD (hopefully). For now, a few recipes:
***
## `SonarWaveletFilteredNoise`
You will need [pytorch_wavelets](https://github.com/fbcotter/pytorch_wavelets) installed in your Python environment to use this one.
Allows filtering another noise source using wavelets. Parameters are specified using YAML (or JSON) in the text widget. The defaults are:
```yaml
use_dtcwt: false
mode: periodization
level: 3
wave: haar
# Only used in DTCWT mode.
qshift: qshift_a
# Only used in DTCWT mode.
biort: near_sym_a
# Additional parameters for the inverse wavelet operation
# are null by default and will use whatever the
# forward parameter is set to:
# inv_mode, inv_wave, inv_biort, inv_qshift
# Note: Using different parameters for the inverse wavelet
# operation may not work well (or just fail entirely).
# Scale for the lowpass filter.
yl_scale: 1.0
# Scales for the highpass filter. Can be a single value (null is basically 1.0).
yh_scales: null
```
***
### `SonarAdvancedDistroNoise`
See: https://pytorch.org/docs/stable/distributions.html
For the most part, we just pass parameters directly to PyTorch's distribution classes. Some of them have specific requirements so it is possible to set invalid parameters.
It may be more convenient to specify parameters using the `SonarCustomNoiseAdv` node than this gigantic monstrosity of a node. **Note**: In that case, pass the distribution name using `distro`, i.e. `distro: laplacian`.
Common parameters:
* `quantile_norm`: When enabled, will normalize generated noise to this quantile (i.e. 0.75 means outliers >75% will be clipped). Set to 1.0 or 0.0 to disable quantile normalization. A value like 0.75 or 0.85 should be reasonable, it really depends on the distribution and how many of the values are extreme. Some actually work better with quantile normalization disabled.
* `quantile_norm_mode`: Controls what dimensions quantile normalization uses. By default, the noise is flattened first. You can try the nonflat versions but they may have a very strong row/column influence. Only applies when quantile_norm is active.
* `result_index`: When noise generation returns a batch of items, it will select the specified index. Negative indexes count from the end. Values outside the valid range will be automatically adjusted. You may enter a space-separated list of values for the case where there might be multiple added batch dimensions. Excess batch dimensions are removed from the end, indexe from result_index are used in order so you may want to enter the indexes in reverse order. Example: If your noise has shape `(1, 4, 3, 3)` and two 2-sized batch dims are added resulting in `(1, 4, 3, 3, 2, 2)` and you wanted index 0 from the first additional batch dimension and 1 from the second you would use result_index: `1 0`
Individual distributions have parameters beginning with their name, i.e. `laplacian_loc`. Parameters that are string inputs usually allow entering multiple space-separated items. This will usually result in the output noise being a batch, which can be selected with the `result_index` parameter.
Suggestions for fun distributions to try: Wishart and VonMises can produce some interesting results.
***
### `SonarModulatedNoise`
Experimental noise modulation based on code stolen from
+117 -14
View File
@@ -1,22 +1,125 @@
from __future__ import annotations
import contextlib
import importlib
import sys
from functools import partial
from typing import TYPE_CHECKING, Callable, NamedTuple
MODULES = {}
if TYPE_CHECKING:
from types import ModuleType
with contextlib.suppress(ImportError, NotImplementedError):
bleh = importlib.import_module("custom_nodes.ComfyUI-bleh")
bleh_version = getattr(bleh, "BLEH_VERSION", -1)
if bleh_version < 1:
raise NotImplementedError
MODULES["bleh"] = bleh
with contextlib.suppress(ImportError, NotImplementedError):
import custom_nodes.ComfyUI_restart_sampling as rs
class Integrations:
class Integration(NamedTuple):
key: str
module_name: str
handler: Callable | None = None
def __init__(self):
self.initialized = False
self.modules = {}
self.init_handlers = []
self.handlers = []
def __getitem__(self, key):
return self.modules[key]
def __contains__(self, key):
return key in self.modules
def __getattr__(self, key):
return self.modules.get(key)
@staticmethod
def get_custom_node(name: str) -> ModuleType | None:
module_key = f"custom_nodes.{name}"
with contextlib.suppress(StopIteration):
spec = importlib.util.find_spec(module_key)
if spec is None:
return None
return next(
v
for v in sys.modules.copy().values()
if hasattr(v, "__spec__")
and v.__spec__ is not None
and v.__spec__.origin == spec.origin
)
return None
def register_init_handler(self, handler):
self.init_handlers.append(handler)
def register_integration(self, key: str, module_name: str, handler=None) -> None:
if self.initialized:
raise ValueError(
"Internal error: Cannot register integration after initialization",
)
if any(item[0] == key or item[1] == module_name for item in self.handlers):
errstr = (
f"Module {module_name} ({key}) already in integration handlers list!"
)
raise ValueError(errstr)
self.handlers.append(self.Integration(key, module_name, handler))
def initialize(self) -> None:
if self.initialized:
return
self.initialized = True
for ih in self.handlers:
module = self.get_custom_node(ih.module_name)
if module is None:
continue
if ih.handler is not None:
module = ih.handler(module)
if module is not None:
self.modules[ih.key] = module
for init_handler in self.init_handlers:
init_handler(self)
class SonarIntegrations(Integrations):
def __init__(self, *args: list, **kwargs: dict):
super().__init__(*args, **kwargs)
self.register_integration("bleh", "ComfyUI-bleh", self.bleh_integration)
self.register_integration(
"restart",
"ComfyUI_restart_sampling",
self.restart_integration,
)
@classmethod
def bleh_integration(cls, module: ModuleType) -> ModuleType | None:
bleh_version = getattr(module, "BLEH_VERSION", -1)
if bleh_version < 1:
return None
return module
@classmethod
def restart_integration(cls, module: ModuleType) -> ModuleType | None:
if hasattr(module, "restart_sampling") and hasattr(
module.restart_sampling,
"DEFAULT_SEGMENTS",
):
return module
return None
MODULES = SonarIntegrations()
class IntegratedNode(type):
@staticmethod
def wrap_INPUT_TYPES(orig_method: Callable, *args: list, **kwargs: dict) -> dict:
MODULES.initialize()
return orig_method(*args, **kwargs)
def __new__(cls: type, name: str, bases: tuple, attrs: dict) -> object:
obj = type.__new__(cls, name, bases, attrs)
if hasattr(obj, "INPUT_TYPES"):
obj.INPUT_TYPES = partial(cls.wrap_INPUT_TYPES, obj.INPUT_TYPES)
return obj
if not hasattr(rs.restart_sampling, "DEFAULT_SEGMENTS"):
# Dumb test but this should only exist in restart sampling versions that
# support plugging in custom noise.
raise NotImplementedError
MODULES["restart"] = rs
__all__ = ("MODULES",)
+7 -13
View File
@@ -2,7 +2,8 @@ from __future__ import annotations
import torch
from .external import MODULES as EXTERNAL_MODULES
from . import utils
from .external import IntegratedNode
from .powernoise import PowerFilter
@@ -28,14 +29,7 @@ def ffilter(x, pfilter, normalization_factor=1.0, cfg_idx=None, filter_cache=Non
return x_filt.to(x.dtype, non_blocking=True)
BLEND_OPS = (
{"lerp": torch.lerp}
if "bleh" not in EXTERNAL_MODULES
else EXTERNAL_MODULES["bleh"].py.latent_utils.BLENDING_MODES
)
class FreeUExtremeConfigNode:
class FreeUExtremeConfigNode(metaclass=IntegratedNode):
DESCRIPTION = "Allows setting configuration for FreeU Extreme."
RETURN_TYPES = ("FRUX_CONFIG",)
FUNCTION = "go"
@@ -150,7 +144,7 @@ class FreeUExtremeConfigNode:
},
),
"blend_mode": (
tuple(BLEND_OPS.keys()),
tuple(utils.BLENDING_MODES.keys()),
{
"tooltip": "Mode used when blending. Generally only has an effect when blend is set to values other than 0 or 1",
},
@@ -302,7 +296,7 @@ class FreeUExtremeConfig:
x[:, slice_offs : slice_offs + slice_size] = (
xslice
if self.blend == 1.0
else BLEND_OPS[self.blend_mode](
else utils.BLENDING_MODES[self.blend_mode](
x[:, slice_offs : slice_offs + slice_size],
xslice,
self.blend,
@@ -331,12 +325,12 @@ class FreeUExtremeConfig:
def clone(self):
return self.__class__(**{k: getattr(self, k) for k in self._keys})
def __repr__(self): # noqa: D105
def __repr__(self):
meh = {k: getattr(self, k) for k in self._keys}
return f"<FRUXConfig: {meh}>"
class FreeUExtremeNode:
class FreeUExtremeNode(metaclass=IntegratedNode):
DESCRIPTION = "Main FreeU Extreme node. Allows patching a model with the FreeU (V2) effect with more control."
RETURN_TYPES = ("MODEL",)
FUNCTION = "go"
+790 -365
View File
File diff suppressed because it is too large Load Diff
+521 -233
View File
@@ -6,18 +6,30 @@ from typing import Callable
import comfy
import torch
import yaml
from comfy.k_diffusion import sampling
from torch import Tensor
from . import external
from . import external, utils
from .noise_generation import *
from .sonar import SonarGuidanceMixin
from .utils import crop_samples, scale_noise
# ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311
class CustomNoiseItemBase(abc.ABC):
def __init__(self, factor, **kwargs):
def __init__(self, factor, *, yaml_parameters=None, **kwargs):
if yaml_parameters:
extra_params = yaml.safe_load(yaml_parameters)
if extra_params is None:
pass
elif not isinstance(extra_params, dict):
raise ValueError(
"CustomNoiseItem: yaml_parameters must either be null or an object",
)
else:
kwargs["ns_kwargs"] = extra_params
self.factor = factor
self.keys = set(kwargs.keys())
for k, v in kwargs.items():
@@ -46,6 +58,7 @@ class CustomNoiseItemBase(abc.ABC):
seed=None,
cpu=True,
normalized=True,
**kwargs,
):
raise NotImplementedError
@@ -65,16 +78,25 @@ class CustomNoiseItem(CustomNoiseItemBase):
seed=None,
cpu=True,
normalized=True,
**kwargs,
):
ns_kwargs = getattr(self, "ns_kwargs", {}).copy()
# print("NS KWARGS", ns_kwargs)
return get_noise_sampler(
self.noise_type,
x,
sigma_min,
sigma_max,
seed=seed,
cpu=cpu,
seed=ns_kwargs.pop("seed", seed),
cpu=ns_kwargs.pop("cpu", cpu),
factor=self.factor,
normalized=self.get_normalize("normalize", normalized),
normalized=ns_kwargs.pop(
"normalized",
self.get_normalize("normalize", normalized),
),
**ns_kwargs,
**kwargs,
)
@@ -155,6 +177,7 @@ class NoiseSampler:
make_noise_sampler: Callable | None = None,
normalized=False,
factor: float = 1.0,
**kwargs,
):
self.factor = factor
self.normalized = normalized
@@ -164,16 +187,18 @@ class NoiseSampler:
try:
self.noise_sampler = make_noise_sampler(
x,
transform(torch.as_tensor(sigma_min))
sigma_min=transform(torch.as_tensor(sigma_min))
if sigma_min is not None
else None,
transform(torch.as_tensor(sigma_max))
sigma_max=transform(torch.as_tensor(sigma_max))
if sigma_max is not None
else None,
seed=seed,
cpu=cpu,
**kwargs,
)
except TypeError as _exc:
print("GOT EXC", _exc)
self.noise_sampler = make_noise_sampler(x)
@classmethod
@@ -216,7 +241,7 @@ class AdvancedNoiseBase(CustomNoiseItemBase):
v = getattr(self, k, None)
if v is not None:
noise_sampler_kwargs[k] = v
self.sampler_factory = NoiseSampler.simple(
self.sampler_factory = NoiseSampler.wrap(
partial(self.ns_factory, **noise_sampler_kwargs),
)
@@ -229,9 +254,9 @@ class AdvancedPyramidNoise(AdvancedNoiseBase):
ns_factory_arg_keys = ("discount", "iterations", "upscale_mode")
pyramid_variants_map = { # noqa: RUF012
"pyramid": pyramid_noise_like,
"pyramid_old": pyramid_old_noise_like,
"highres_pyramid": highres_pyramid_noise_like,
"pyramid": PyramidNoiseGenerator,
"pyramid_old": PyramidOldNoiseGenerator,
"highres_pyramid": HighresPyramidNoiseGenerator,
}
@property
@@ -244,15 +269,33 @@ class Advanced1fNoise(AdvancedNoiseBase):
@property
def ns_factory(self):
return onef_noise_like
return OneFNoiseGenerator
class AdvancedPowerLawNoise(AdvancedNoiseBase):
ns_factory_arg_keys = ("alpha", "div_max_dims", "use_sign")
@classmethod
@property
def ns_factory(self):
return powerlaw_noise_like
def ns_factory(cls):
return PowerLawNoiseGenerator
class AdvancedDistroNoise(AdvancedNoiseBase):
distro_params = DistroNoiseGenerator.build_params()
ns_factory_arg_keys = (
"distro",
"quantile_norm",
"quantile_norm_dim",
"quantile_norm_flatten",
"result_index",
*distro_params.keys(),
)
@classmethod
@property
def ns_factory(cls):
return DistroNoiseGenerator
class CompositeNoise(CustomNoiseItemBase):
@@ -446,6 +489,10 @@ class ScheduledNoise(CustomNoiseItemBase):
return torch.zeros_like(x)
def noise_sampler(s, sn):
if s is None or sn is None:
raise ValueError(
"ScheduledNoise requires sigma, sigma_next to be passed",
)
noise = (ns if end_sigma <= s <= start_sigma else nsa)(s, sn)
return scale_noise(noise, factor, normalized=normalize)
@@ -997,271 +1044,510 @@ class BlendedNoise(CustomNoiseItemBase):
return noise_sampler
if "bleh" in external.MODULES:
bleh = external.MODULES["bleh"]
BLU = bleh.py.latent_utils
BOPS = bleh.py.nodes.ops
class BlendFilterNoise(CustomNoiseItemBase):
def __init__(
self,
class ResizedNoise(CustomNoiseItemBase):
def __init__(
self,
factor,
*,
width,
height,
downscale_strategy,
initial_reference,
crop_offset_horizontal,
crop_offset_vertical,
crop_mode,
upscale_mode,
downscale_mode,
normalize,
custom_noise,
):
if len(custom_noise.items) == 0:
raise ValueError("ResizedNoise requires at least one noise item")
super().__init__(
factor,
*,
noise,
blend_mode,
ffilter,
ffilter_scale,
ffilter_strength,
ffilter_threshold,
enhance_mode,
enhance_strength,
affect,
normalize_result,
normalize_noise,
):
if len(noise.items) == 0:
raise ValueError("BlendFilterNoise requires at least one noise item")
super().__init__(
factor,
noise=noise.clone(),
blend_mode=blend_mode,
ffilter=ffilter,
ffilter_scale=ffilter_scale,
ffilter_strength=ffilter_strength,
ffilter_threshold=ffilter_threshold,
enhance_mode=enhance_mode,
enhance_strength=enhance_strength,
affect=affect,
normalize_result=normalize_result,
normalize_noise=normalize_noise,
)
width=width,
height=height,
downscale_strategy=downscale_strategy,
initial_reference=initial_reference,
crop_offset_horizontal=crop_offset_horizontal,
crop_offset_vertical=crop_offset_vertical,
crop_mode=crop_mode,
upscale_mode=upscale_mode,
downscale_mode=downscale_mode,
custom_noise=custom_noise.clone(),
normalize=normalize,
)
def clone_key(self, k):
if k == "noise":
return self.noise.clone()
return super().clone_key(k)
def clone_key(self, k):
if k == "custom_noise":
return self.custom_noise.clone()
return super().clone_key(k)
def apply_effects(self, noise, sigma):
if self.ffilter:
noise = BLU.ffilter(
noise,
self.ffilter_threshold,
self.ffilter_scale,
self.ffilter,
self.ffilter_strength,
)
if self.enhance_mode != "none" and self.enhance_strength != 0:
noise = BLU.enhance_tensor(
noise,
self.enhance_mode,
self.enhance_strength,
sigma=sigma,
skip_multiplier=0,
adjust_scale=False,
)
return noise
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
factor = self.factor
noise_items = self.noise.items
noise_samplers = tuple(
ni.make_noise_sampler(x, *args, normalized=False, **kwargs)
for ni in noise_items
)
num_samplers = len(noise_samplers)
normalize_noise = self.get_normalize(
"normalize_noise",
normalized or num_samplers > 1,
)
normalize_result = self.get_normalize("normalize_result", normalized)
noise_effects = self.affect in {"noise", "both"}
result_effects = self.affect in {"result", "both"}
noise_init = torch.zeros_like(x)
def noise_sampler(s, sn):
noise = noise_init.clone()
for ni, ns in zip(noise_items, noise_samplers):
curr_noise = scale_noise(ns(s, sn), normalized=normalize_noise)
if noise_effects:
curr_noise = self.apply_effects(curr_noise, s)
if self.blend_mode == "simple_add":
noise += curr_noise.mul_(ni.factor)
else:
noise = BLU.BLENDING_MODES[self.blend_mode](
noise,
curr_noise,
ni.factor,
)
del curr_noise
noise = scale_noise(noise, factor, normalized=normalize_result)
if result_effects:
noise = self.apply_effects(noise, s)
return noise
return noise_sampler
class BlehOpsNoise(CustomNoiseItemBase):
def __init__(
self,
factor,
*,
noise,
rules,
normalize,
):
if len(noise.items) == 0:
raise ValueError("BlehOpsNoise requires at least one noise item")
super().__init__(
factor,
noise=noise.clone(),
rules=rules,
normalize=normalize,
)
def clone_key(self, k):
if k == "noise":
return self.noise.clone()
return super().clone_key(k)
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
factor = self.factor
normalize = self.get_normalize("normalize", normalized)
rulegroup = self.rules
internal_ns = self.noise.make_noise_sampler(
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
if x.ndim < 3:
raise ValueError("ResizedNoise can only handle 3+ dimensional latents")
factor = self.factor
normalize = self.get_normalize("normalize", normalized)
xh, xw = x.shape[-2:]
nh, nw = self.height // 8, self.width // 8
offsh, offsw = self.crop_offset_vertical // 8, self.crop_offset_horizontal // 8
if xh == nh and xw == nw:
ns = self.custom_noise.make_noise_sampler(
x,
*args,
normalized=False,
normalized=normalize,
**kwargs,
)
def noise_sampler(s, sn):
noise = internal_ns(s, sn)
if len(rulegroup.rules):
state = {
BOPS.CondType.TYPE: BOPS.PatchType.LATENT,
BOPS.CondType.PERCENT: 0.0,
BOPS.CondType.BLOCK: -1,
BOPS.CondType.STAGE: -1,
"sigma": None if s is None else s,
"h": noise,
"hsp": x.detach().clone(),
"target": "h",
}
noise = rulegroup.eval(state, toplevel=True)["h"]
return scale_noise(noise, factor, normalized=normalize)
def noise_sampler(*args, **kwargs):
return ns(*args, **kwargs).mul_(factor)
return noise_sampler
upscale_mode = self.upscale_mode
downscale_mode = self.downscale_mode
crop_mode = self.crop_mode
x_all_bigger = xh >= nh and xw >= nw
x_any_bigger = xh >= nh or xw >= nw
if x_all_bigger:
if self.initial_reference == "prefer_crop":
x = crop_samples(
x,
nw,
nh,
mode=self.crop_mode,
offset_width=offsw,
offset_height=offsh,
)
else:
x = utils.scale_samples(x, nw, nh, mode=self.downscale_mode)
output = partial(
utils.scale_samples,
width=xw,
height=xh,
mode=upscale_mode,
)
else:
x = utils.scale_samples(x, nw, nh, mode=self.upscale_mode)
if x_any_bigger:
output = partial(
utils.scale_samples,
width=xw,
height=xh,
mode=upscale_mode,
)
elif self.downscale_strategy == "scale":
output = partial(
utils.scale_samples,
width=xw,
height=xh,
mode=downscale_mode,
)
else:
output = partial(
crop_samples,
width=xw,
height=xh,
mode=crop_mode,
offset_width=offsw,
offset_height=offsh,
)
ns = self.custom_noise.make_noise_sampler(x, *args, normalized=False, **kwargs)
del x
def noise_sampler(*args, **kwargs):
return output(
scale_noise(ns(*args, **kwargs), factor, normalized=normalize),
)
return noise_sampler
class WaveletFilteredNoise(CustomNoiseItemBase):
def __init__(self, factor, *, normalize, noise, normalize_noise=False, **kwargs):
super().__init__(
factor,
noise=noise,
normalize=normalize,
normalize_noise=normalize_noise,
**kwargs,
)
def clone_key(self, k):
if k == "noise":
return self.noise.clone()
return super().clone_key(k)
def make_noise_sampler(
self,
x,
sigma_min,
sigma_max,
*args,
normalized=True,
**kwargs,
):
factor = self.factor
normalize = self.get_normalize("normalize", normalized)
internal_ns = self.noise.make_noise_sampler(
x,
*args,
sigma_min=sigma_min,
sigma_max=sigma_max,
normalized=self.normalize_noise,
**kwargs,
)
ns_kwargs = getattr(self, "ns_kwargs", {}).copy()
# print("WF:NS KWARGS", ns_kwargs)
kwargs |= ns_kwargs
ns = WaveletNoiseGenerator(
x,
*args,
sigma_min=sigma_min,
sigma_max=sigma_max,
normalized=False,
noise_sampler=internal_ns,
**kwargs,
)
def noise_sampler(sigma, sigma_next):
return scale_noise(
ns(sigma, sigma_next),
factor,
normalized=normalize,
)
return noise_sampler
class BlendFilterNoise(CustomNoiseItemBase):
def __init__(
self,
factor,
*,
noise,
blend_mode,
ffilter,
ffilter_scale,
ffilter_strength,
ffilter_threshold,
enhance_mode,
enhance_strength,
affect,
normalize_result,
normalize_noise,
):
if len(noise.items) == 0:
raise ValueError("BlendFilterNoise requires at least one noise item")
super().__init__(
factor,
noise=noise.clone(),
blend_mode=blend_mode,
ffilter=ffilter,
ffilter_scale=ffilter_scale,
ffilter_strength=ffilter_strength,
ffilter_threshold=ffilter_threshold,
enhance_mode=enhance_mode,
enhance_strength=enhance_strength,
affect=affect,
normalize_result=normalize_result,
normalize_noise=normalize_noise,
)
def clone_key(self, k):
if k == "noise":
return self.noise.clone()
return super().clone_key(k)
def apply_effects(self, noise, sigma):
blu = external.MODULES.bleh.py.latent_utils
if self.ffilter:
noise = blu.ffilter(
noise,
self.ffilter_threshold,
self.ffilter_scale,
self.ffilter,
self.ffilter_strength,
)
if self.enhance_mode != "none" and self.enhance_strength != 0:
noise = blu.enhance_tensor(
noise,
self.enhance_mode,
self.enhance_strength,
sigma=sigma,
skip_multiplier=0,
adjust_scale=False,
)
return noise
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
factor = self.factor
noise_items = self.noise.items
noise_samplers = tuple(
ni.make_noise_sampler(x, *args, normalized=False, **kwargs)
for ni in noise_items
)
num_samplers = len(noise_samplers)
normalize_noise = self.get_normalize(
"normalize_noise",
normalized or num_samplers > 1,
)
normalize_result = self.get_normalize("normalize_result", normalized)
noise_effects = self.affect in {"noise", "both"}
result_effects = self.affect in {"result", "both"}
noise_init = torch.zeros_like(x)
def noise_sampler(s, sn):
noise = noise_init.clone()
for ni, ns in zip(noise_items, noise_samplers):
curr_noise = scale_noise(ns(s, sn), normalized=normalize_noise)
if noise_effects:
curr_noise = self.apply_effects(curr_noise, s)
if self.blend_mode == "simple_add":
noise += curr_noise.mul_(ni.factor)
else:
noise = utils.BLENDING_MODES[self.blend_mode](
noise,
curr_noise,
ni.factor,
)
del curr_noise
noise = scale_noise(noise, factor, normalized=normalize_result)
if result_effects:
noise = self.apply_effects(noise, s)
return noise
return noise_sampler
class BlehOpsNoise(CustomNoiseItemBase):
def __init__(
self,
factor,
*,
noise,
rules,
normalize,
):
if len(noise.items) == 0:
raise ValueError("BlehOpsNoise requires at least one noise item")
super().__init__(
factor,
noise=noise.clone(),
rules=rules,
normalize=normalize,
)
def clone_key(self, k):
if k == "noise":
return self.noise.clone()
return super().clone_key(k)
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
bops = external.MODULES.bleh.py.nodes.ops
factor = self.factor
normalize = self.get_normalize("normalize", normalized)
rulegroup = self.rules
internal_ns = self.noise.make_noise_sampler(
x,
*args,
normalized=False,
**kwargs,
)
def noise_sampler(s, sn):
noise = internal_ns(s, sn)
if len(rulegroup.rules):
state = {
bops.CondType.TYPE: bops.PatchType.LATENT,
bops.CondType.PERCENT: 0.0,
bops.CondType.BLOCK: -1,
bops.CondType.STAGE: -1,
"sigma": None if s is None else s,
"h": noise,
"hsp": x.detach().clone(),
"target": "h",
}
noise = rulegroup.eval(state, toplevel=True)["h"]
return scale_noise(noise, factor, normalized=normalize)
return noise_sampler
NOISE_SAMPLERS: dict[NoiseType, Callable] = {
NoiseType.BROWNIAN: NoiseSampler.wrap(sampling.BrownianTreeNoiseSampler),
NoiseType.GAUSSIAN: NoiseSampler.simple(torch.randn_like),
NoiseType.UNIFORM: NoiseSampler.simple(uniform_noise_like),
NoiseType.PERLIN: NoiseSampler.simple(rand_perlin_like),
NoiseType.STUDENTT: NoiseSampler.simple(studentt_noise_like),
NoiseType.ONEF_PINKISH: NoiseSampler.simple(partial(onef_noise_like, alpha=-0.5)),
NoiseType.ONEF_GREENISH: NoiseSampler.simple(partial(onef_noise_like, alpha=0.5)),
NoiseType.ONEF_PINKISHGREENISH: NoiseSampler.simple(
lambda x: onef_noise_like(x, alpha=0.5)
.add_(onef_noise_like(x, alpha=-0.5))
.mul_(0.5),
),
NoiseType.ONEF_PINKISH_MIX: NoiseSampler.simple(
lambda x: onef_noise_like(x, alpha=-0.5)
.mul_(-1.0)
.add_(onef_noise_like(x, alpha=-0.5))
.mul_(0.5),
),
NoiseType.ONEF_GREENISH_MIX: NoiseSampler.simple(
lambda x: onef_noise_like(x, alpha=0.5)
.mul_(-1.0)
.add_(onef_noise_like(x, alpha=0.5))
.mul_(0.5),
),
NoiseType.WHITE: NoiseSampler.simple(
NoiseType.BROWNIAN: NoiseSampler.wrap(BrownianNoiseGenerator),
NoiseType.DISTRO: NoiseSampler.wrap(DistroNoiseGenerator),
NoiseType.GAUSSIAN: NoiseSampler.wrap(GaussianNoiseGenerator),
NoiseType.UNIFORM: NoiseSampler.wrap(UniformNoiseGenerator),
NoiseType.PERLIN: NoiseSampler.wrap(PerlinOldNoiseGenerator),
NoiseType.STUDENTT: NoiseSampler.wrap(StudentTNoiseGenerator),
NoiseType.ONEF_PINKISH: NoiseSampler.wrap(partial(OneFNoiseGenerator, alpha=-0.5)),
NoiseType.ONEF_GREENISH: NoiseSampler.wrap(partial(OneFNoiseGenerator, alpha=0.5)),
NoiseType.ONEF_PINKISHGREENISH: NoiseSampler.wrap(
partial(
powerlaw_noise_like,
MixedNoiseGenerator,
name="onef_pinkishgreenish",
noise_mix=(
(OneFNoiseGenerator, {"alpha": 0.5}, None),
(OneFNoiseGenerator, {"alpha": -0.5}, None),
),
output_fun=lambda t: t.mul_(0.5),
),
),
NoiseType.ONEF_PINKISH_MIX: NoiseSampler.wrap(
partial(
MixedNoiseGenerator,
name="onef_pinkish_mix",
noise_mix=(
(OneFNoiseGenerator, {"alpha": 0.5}, lambda t: t.mul_(-1.0)),
(OneFNoiseGenerator, {"alpha": 0.5}, None),
),
output_fun=lambda t: t.mul_(0.5),
),
),
NoiseType.ONEF_GREENISH_MIX: NoiseSampler.wrap(
partial(
MixedNoiseGenerator,
name="onef_greenish_mix",
noise_mix=(
(OneFNoiseGenerator, {"alpha": 0.5}, lambda t: t.mul_(-1.0)),
(OneFNoiseGenerator, {"alpha": 0.5}, None),
),
output_fun=lambda t: t.mul_(0.5),
),
),
NoiseType.WHITE: NoiseSampler.wrap(
partial(
PowerLawNoiseGenerator,
alpha=0.0,
use_sign=True,
),
),
NoiseType.GREY: NoiseSampler.simple(
NoiseType.GREY: NoiseSampler.wrap(
partial(
powerlaw_noise_like,
PowerLawNoiseGenerator,
alpha=0.0,
use_sign=False,
),
),
NoiseType.VELVET: NoiseSampler.simple(
NoiseType.VELVET: NoiseSampler.wrap(
partial(
powerlaw_noise_like,
PowerLawNoiseGenerator,
alpha=1.0,
use_sign=True,
div_max_dims=(-3, -2, -1),
),
),
NoiseType.VIOLET: NoiseSampler.simple(
NoiseType.VIOLET: NoiseSampler.wrap(
partial(
powerlaw_noise_like,
PowerLawNoiseGenerator,
alpha=0.5,
use_sign=True,
div_max_dims=(-3, -2, -1),
),
),
NoiseType.PINK_OLD: NoiseSampler.simple(pink_noise_old_like),
NoiseType.HIGHRES_PYRAMID: NoiseSampler.simple(highres_pyramid_noise_like),
NoiseType.PYRAMID: NoiseSampler.simple(pyramid_noise_like),
NoiseType.RAINBOW_MILD: NoiseSampler.simple(
lambda x: green_noise_like(x)
.mul_(0.55)
.add_(rand_perlin_like(x).mul_(0.7))
.mul_(1.15),
NoiseType.WAVELET: NoiseSampler.wrap(WaveletNoiseGenerator),
NoiseType.PINK_OLD: NoiseSampler.wrap(PinkOldNoiseGenerator),
NoiseType.HIGHRES_PYRAMID: NoiseSampler.wrap(HighresPyramidNoiseGenerator),
NoiseType.PYRAMID: NoiseSampler.wrap(PyramidNoiseGenerator),
NoiseType.RAINBOW_MILD: NoiseSampler.wrap(
partial(
MixedNoiseGenerator,
name="rainbow_mild",
noise_mix=(
(GreenTestNoiseGenerator, {}, lambda t: t.mul_(0.55)),
(GreenTestNoiseGenerator, {}, lambda t: t.mul_(0.7)),
),
output_fun=lambda t: t.mul_(1.15),
),
),
NoiseType.RAINBOW_INTENSE: NoiseSampler.simple(
lambda x: green_noise_like(x)
.mul_(0.75)
.add_(rand_perlin_like(x).mul_(0.5))
.mul_(1.15),
NoiseType.RAINBOW_INTENSE: NoiseSampler.wrap(
partial(
MixedNoiseGenerator,
name="rainbow_intense",
noise_mix=(
(GreenTestNoiseGenerator, {}, lambda t: t.mul_(0.75)),
(GreenTestNoiseGenerator, {}, lambda t: t.mul_(0.5)),
),
output_fun=lambda t: t.mul_(1.15),
),
),
NoiseType.LAPLACIAN: NoiseSampler.simple(laplacian_noise_like),
NoiseType.POWER_OLD: NoiseSampler.simple(power_noise_old_like),
NoiseType.GREEN_TEST: NoiseSampler.simple(green_noise_like),
NoiseType.PYRAMID_OLD: NoiseSampler.simple(pyramid_old_noise_like),
NoiseType.PYRAMID_BISLERP: NoiseSampler.simple(
partial(pyramid_noise_like, upscale_mode="bislerp"),
NoiseType.LAPLACIAN: NoiseSampler.wrap(LaplacianNoiseGenerator),
NoiseType.POWER_OLD: NoiseSampler.wrap(PowerOldNoiseGenerator),
NoiseType.GREEN_TEST: NoiseSampler.wrap(GreenTestNoiseGenerator),
NoiseType.PYRAMID_OLD: NoiseSampler.wrap(PyramidOldNoiseGenerator),
NoiseType.PYRAMID_BISLERP: NoiseSampler.wrap(
partial(PyramidNoiseGenerator, upscale_mode="bislerp"),
),
NoiseType.HIGHRES_PYRAMID_BISLERP: NoiseSampler.simple(
partial(highres_pyramid_noise_like, upscale_mode="bislerp"),
NoiseType.HIGHRES_PYRAMID_BISLERP: NoiseSampler.wrap(
partial(HighresPyramidNoiseGenerator, upscale_mode="bislerp"),
),
NoiseType.PYRAMID_AREA: NoiseSampler.simple(
partial(pyramid_noise_like, upscale_mode="area"),
NoiseType.PYRAMID_AREA: NoiseSampler.wrap(
partial(PyramidNoiseGenerator, upscale_mode="area"),
),
NoiseType.HIGHRES_PYRAMID_AREA: NoiseSampler.simple(
partial(highres_pyramid_noise_like, upscale_mode="area"),
NoiseType.HIGHRES_PYRAMID_AREA: NoiseSampler.wrap(
partial(HighresPyramidNoiseGenerator, upscale_mode="area"),
),
NoiseType.PYRAMID_OLD_BISLERP: NoiseSampler.simple(
partial(pyramid_old_noise_like, upscale_mode="bislerp"),
NoiseType.PYRAMID_OLD_BISLERP: NoiseSampler.wrap(
partial(PyramidOldNoiseGenerator, upscale_mode="bislerp"),
),
NoiseType.PYRAMID_OLD_AREA: NoiseSampler.simple(
partial(pyramid_old_noise_like, upscale_mode="area"),
NoiseType.PYRAMID_OLD_AREA: NoiseSampler.wrap(
partial(PyramidOldNoiseGenerator, upscale_mode="area"),
),
NoiseType.PYRAMID_DISCOUNT5: NoiseSampler.simple(
partial(pyramid_noise_like, discount=0.5),
NoiseType.PYRAMID_DISCOUNT5: NoiseSampler.wrap(
partial(PyramidNoiseGenerator, discount=0.5),
),
NoiseType.PYRAMID_MIX: NoiseSampler.simple(
lambda x: pyramid_noise_like(x, discount=0.6)
.mul_(0.2)
.add_(pyramid_noise_like(x, discount=0.6).mul_(-0.8)),
NoiseType.PYRAMID_MIX: NoiseSampler.wrap(
partial(
MixedNoiseGenerator,
name="pyramid_mix",
noise_mix=(
(PyramidNoiseGenerator, {"discount": 0.6}, lambda t: t.mul_(0.2)),
(PyramidNoiseGenerator, {"discount": 0.6}, lambda t: t.mul_(-0.8)),
),
),
),
NoiseType.PYRAMID_MIX_AREA: NoiseSampler.simple(
lambda x: pyramid_noise_like(x, discount=0.5, upscale_mode="area")
.mul_(0.2)
.add_(pyramid_noise_like(x, discount=0.5, upscale_mode="area").mul_(-0.8)),
NoiseType.PYRAMID_MIX_AREA: NoiseSampler.wrap(
partial(
MixedNoiseGenerator,
name="pyramid_mix_area",
noise_mix=(
(
PyramidNoiseGenerator,
{"discount": 0.5, "upscale_mode": "area"},
lambda t: t.mul_(0.2),
),
(
PyramidNoiseGenerator,
{"discount": 0.5, "upscale_mode": "area"},
lambda t: t.mul_(-0.8),
),
),
),
),
NoiseType.PYRAMID_MIX_BISLERP: NoiseSampler.simple(
lambda x: pyramid_noise_like(x, discount=0.5, upscale_mode="bislerp")
.mul_(0.2)
.add_(pyramid_noise_like(x, discount=0.5, upscale_mode="bislerp").mul_(-0.8)),
NoiseType.PYRAMID_MIX_BISLERP: NoiseSampler.wrap(
partial(
MixedNoiseGenerator,
name="pyramid_mix_bislerp",
noise_mix=(
(
PyramidNoiseGenerator,
{
"discount": 0.5,
"upscale_mode": "bislerp",
},
lambda t: t.mul_(0.2),
),
(
PyramidNoiseGenerator,
{
"discount": 0.5,
"upscale_mode": "bislerp",
},
lambda t: t.mul_(-0.8),
),
),
),
),
}
@@ -1275,6 +1561,7 @@ def get_noise_sampler(
cpu: bool = True,
factor: float = 1.0,
normalized=False,
**kwargs,
) -> Callable:
if noise_type is None:
noise_type = NoiseType.GAUSSIAN
@@ -1293,4 +1580,5 @@ def get_noise_sampler(
cpu=cpu,
factor=factor,
normalized=normalized,
**kwargs,
)
+1266 -441
View File
File diff suppressed because it is too large Load Diff
+4 -3
View File
@@ -17,12 +17,13 @@ from PIL import Image
from torch import Tensor
from .nodes import (
NOISE_INPUT_TYPES_HINT,
WILDCARD_NOISE,
SonarCustomNoiseNodeBase,
SonarNormalizeNoiseNodeMixin,
)
from .noise import CustomNoiseItemBase
from .noise_generation import scale_noise
from .utils import scale_noise
# ruff: noqa: ANN003, FBT001, FBT002
@@ -635,7 +636,7 @@ class SonarPowerNoiseNode(SonarCustomNoiseNodeBase):
"max": 100.0,
"step": 0.001,
"round": False,
"tooltip": "Attempts to desaturate thelatent by injecting the average across channels (controlled by channel_correction). Applied after mix.",
"tooltip": "Attempts to desaturate the latent by injecting the average across channels (controlled by channel_correction). Applied after mix.",
},
),
"channel_correlation": (
@@ -694,7 +695,7 @@ class SonarPowerFilterNoiseNode(SonarPowerNoiseNode, SonarNormalizeNoiseNodeMixi
"sonar_custom_noise": (
WILDCARD_NOISE,
{
"tooltip": "Custom noise type to filter.",
"tooltip": f"Custom noise type to filter.\n{NOISE_INPUT_TYPES_HINT}",
},
),
"sonar_power_filter": (
+398 -189
View File
@@ -4,22 +4,24 @@ from __future__ import annotations
import importlib
from enum import Enum, auto
from functools import lru_cache
from sys import stderr
from typing import Any, Callable, NamedTuple
import torch
from comfy.k_diffusion import sampling
from comfy.k_diffusion.sampling import get_ancestral_step, to_d
from comfy.samplers import KSampler, k_diffusion_sampling
from torch import Tensor
from tqdm.auto import trange
from . import noise
from . import noise, utils
class HistoryType(Enum):
ZERO = auto()
RAND = auto()
SAMPLE = auto()
SAMPLE_NORM = auto()
class GuidanceType(Enum):
@@ -35,15 +37,34 @@ class GuidanceConfig(NamedTuple):
latent: Tensor | None = None
class MomentumMode(Enum):
CLASSIC = auto()
NEW = auto()
DENOISED = auto()
class SonarConfig(NamedTuple):
momentum: float = 0.95
momentum_hist: float = 0.75
direction: float = 1.0
momentum_start_step: int = 0
momentum_end_step: int = 9999
always_update_history: bool = True
momentum_mode: MomentumMode = MomentumMode.NEW
init: HistoryType = HistoryType.ZERO
noise_type: noise.NoiseType | None = None
custom_noise: noise.CustomNoise | None = None
rand_init_noise_type: noise.NoiseType | None = None
rand_init_noise_multiplier: float | int = 1.0
guidance: GuidanceConfig | None = None
blend_mode: str = "lerp"
momentum_blend_mode: str | None = None
history_blend_mode: str | None = None
guidance_blend_mode: str | None = None
def get_with_default(self, k: str, default: Any) -> Any: # noqa: ANN401
val = getattr(self, k)
return val if val is not None else default
class SonarBase:
@@ -53,14 +74,69 @@ class SonarBase:
self.history_d = None
self.cfg = cfg
self.noise_sampler = None
blend_mode = cfg.blend_mode
momentum_blend_mode = cfg.get_with_default("momentum_blend_mode", blend_mode)
history_blend_mode = cfg.get_with_default("history_blend_mode", blend_mode)
guidance_blend_mode = cfg.get_with_default("guidance_blend_mode", blend_mode)
bf = self.blend = utils.BLENDING_MODES[blend_mode]
self.momentum_blend = (
bf
if momentum_blend_mode == blend_mode
else utils.BLENDING_MODES[momentum_blend_mode]
)
self.history_blend = (
bf
if history_blend_mode == blend_mode
else utils.BLENDING_MODES[history_blend_mode]
)
self.guidance_blend = (
bf
if guidance_blend_mode == blend_mode
else utils.BLENDING_MODES[guidance_blend_mode]
)
_cfg_fixups = (
("momentum_mode", MomentumMode),
("init", HistoryType),
("noise_type", noise.NoiseType),
)
@classmethod
def get_config(
cls,
cfg: SonarConfig | None = None,
ext: dict | None = None,
) -> SonarConfig:
cfgdict = ext.copy() if ext is not None else {}
empty = object()
for k, enum_class in cls._cfg_fixups:
val = cfgdict.get(k, empty)
if val is empty:
continue
if isinstance(val, str):
val = getattr(enum_class, val.strip().upper(), empty)
if val is empty:
validstr = ", ".join(enum_class.__members__.keys())
errstr = f"Bad value for {k} of type enum {enum_class.__name__}, must be one of the following: {validstr}"
raise ValueError(errstr)
cfgdict[k] = val
continue
if not isinstance(val, enum_class):
errstr = f"Bad parameter type for {k}: Must be valid string or instance of {enum_class.__name__}"
raise TypeError(errstr)
if cfg is None:
return SonarConfig(**cfgdict)
cfgdict = cfg._asdict() | cfgdict
return SonarConfig(**cfgdict)
def set_noise_sampler(
self,
x: Tensor,
sigmas,
sigmas: Tensor,
noise_sampler: Callable | None,
seed: int | None = None,
):
) -> Callable:
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
if noise_sampler is not None and self.cfg.noise_type not in {
None,
@@ -90,17 +166,32 @@ class SonarBase:
self.noise_sampler = noise_sampler
return noise_sampler
def init_hist_d(self, x: Tensor) -> None:
if self.history_d is not None:
def init_hist_d(
self,
x: Tensor,
denoised: Tensor,
sigma: Tensor,
*,
step: int,
) -> None:
if self.history_d is not None or not self.check_step(step, is_history=True):
return
cfg = self.cfg
init = cfg.init
# memorize delta momentum
if self.cfg.init == HistoryType.ZERO:
self.history_d = 0
elif self.cfg.init == HistoryType.SAMPLE:
self.history_d = x
elif self.cfg.init == HistoryType.RAND:
if init == HistoryType.ZERO:
self.history_d = None
elif init == HistoryType.SAMPLE:
self.history_d = (
x if cfg.momentum_mode != MomentumMode.DENOISED else denoised
)
elif init == HistoryType.SAMPLE_NORM:
self.history_d = (
x if cfg.momentum_mode != MomentumMode.DENOISED else denoised
) / sigma
elif init == HistoryType.RAND:
ns = noise.get_noise_sampler(
self.cfg.rand_init_noise_type,
cfg.rand_init_noise_type,
x,
None,
None,
@@ -109,31 +200,124 @@ class SonarBase:
normalized=True,
)
self.history_d = ns(None, None)
if cfg.rand_init_noise_multiplier != 1:
self.history_d *= cfg.rand_init_noise_multiplier
else:
raise ValueError("Sonar sampler: bad history type")
def update_hist(self, momentum_d):
q = 1.0 - self.cfg.momentum_hist
@property
@lru_cache(maxsize=1) # noqa: B019
def history_ratios(self):
direction = self.cfg.direction
momentum_hist = self.cfg.momentum_hist
return (
momentum_hist,
1.0 + abs(direction) * (1 - momentum_hist)
if direction < 0
else 2.0 - direction,
direction,
)
def check_step(self, step: int, *, is_history: bool = False):
cfg = self.cfg
if is_history and cfg.always_update_history:
return True
return cfg.momentum_start_step <= step <= cfg.momentum_end_step
def update_hist(self, momentum_d: torch.Tensor, step: int) -> None:
hd, cfg = self.history_d, self.cfg
if cfg.momentum_hist == 1 or not self.check_step(step, is_history=True):
return
hd_ratio, hd_scale, md_scale = self.history_ratios
self.history_d = (
momentum_d
if hd is None
else self.history_blend(momentum_d * md_scale, hd * hd_scale, hd_ratio)
)
def momentum_mix(
self,
history: Tensor | None,
item: Tensor,
sigma: Tensor,
*,
is_denoised: bool = False,
momentum=None,
) -> Tensor:
momentum = self.cfg.momentum if momentum is None else momentum
mode = self.cfg.momentum_mode
if (
momentum == 1 # noqa: PLR0916
or history is None
or (mode == MomentumMode.DENOISED and not is_denoised)
or (mode != MomentumMode.DENOISED and is_denoised)
):
return item
return self.momentum_blend(
history * sigma if is_denoised else history,
item,
momentum,
)
def get_momentum_denoised(
self,
x: Tensor,
denoised: Tensor,
sigma: Tensor,
*,
step: int,
momentum: float | None = None,
update_history=True,
) -> Tensor:
hd = self.history_d
if isinstance(hd, int) and hd == 0:
self.history_d = momentum_d
else:
self.history_d = (1.0 - q) * hd + q * momentum_d
momentum_denoised = self.momentum_mix(
hd,
denoised,
sigma,
is_denoised=True,
momentum=momentum,
)
if update_history:
self.init_hist_d(x, denoised, sigma, step=step)
self.update_hist(denoised / sigma, step=step)
return momentum_denoised if self.check_step(step) else denoised
def momentum_step(self, x: Tensor, d: Tensor, dt: Tensor):
if self.cfg.momentum == 1.0:
return x + d * dt
def get_momentum_d(
self,
x: Tensor,
denoised: Tensor,
sigma: Tensor,
*,
step: int,
momentum: float | None = None,
d: Tensor | None = None,
update_history=True,
) -> Tensor:
hd = self.history_d
# correct current `d` with momentum
p = (1.0 - self.cfg.momentum) * self.cfg.direction
momentum_d = (1.0 - p) * d + p * hd
cfg = self.cfg
momentum = cfg.momentum if momentum is None else momentum
mode = cfg.momentum_mode
d = to_d(x, sigma, denoised) if d is None else d
if momentum == 1 or mode == MomentumMode.DENOISED:
return d
momentum_d = self.momentum_mix(hd, d, sigma)
if update_history:
self.init_hist_d(x, denoised, sigma, step=step)
self.update_hist(d if mode == MomentumMode.NEW else momentum_d, step=step)
return momentum_d if self.check_step(step) else d
# Euler method with momentum
x = x + momentum_d * dt # noqa: PLR6104
self.update_hist(momentum_d)
return x
def momentum_step(
self,
step: int,
x: Tensor,
denoised: Tensor,
sigma: Tensor,
sigma_down: Tensor,
) -> Tensor:
dt = sigma_down - sigma
denoised = self.get_momentum_denoised(x, denoised, sigma, step=step)
momentum_d = self.get_momentum_d(x, denoised, sigma, step=step)
return (momentum_d * dt).add_(x)
class SonarGuidanceMixin:
@@ -152,17 +336,26 @@ class SonarGuidanceMixin:
def prepare_ref_latent(latent: Tensor | None) -> Tensor:
if latent is None:
return None
avg_s = latent.mean(dim=[2, 3], keepdim=True)
std_s = latent.std(dim=[2, 3], keepdim=True)
return ((latent - avg_s) / std_s).to(latent.dtype)
avg_s = latent.mean(dim=(-2, -1), keepdim=True)
std_s = latent.std(dim=(-2, -1), keepdim=True)
return (latent - avg_s).div_(std_s).to(latent.dtype)
def guidance_step(self, step_index: int, x: Tensor, denoised: Tensor):
if self.guidance is None or self.guidance.factor == 0.0 or not self.guidance.start_step <= step_index <= self.guidance.end_step:
def guidance_step(self, step_index: int, x: Tensor, denoised: Tensor) -> Tensor:
if (
self.guidance is None
or self.guidance.factor == 0.0
or not self.guidance.start_step <= step_index <= self.guidance.end_step
):
return x
if self.ref_latent.device != x.device:
self.ref_latent = self.ref_latent.to(device=x.device)
if self.guidance.guidance_type == GuidanceType.LINEAR:
return self.guidance_linear(x, self.ref_latent, self.guidance.factor)
return self.guidance_linear(
x,
self.ref_latent,
self.guidance.factor,
blend=self.guidance_blend,
)
if self.guidance.guidance_type == GuidanceType.EULER:
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
return self.guidance_euler(
@@ -184,20 +377,26 @@ class SonarGuidanceMixin:
ref_latent: Tensor,
factor: float = 0.2,
) -> Tensor:
avg_t = denoised.mean(dim=[1, 2, 3], keepdim=True)
std_t = denoised.std(dim=[1, 2, 3], keepdim=True)
avg_t = denoised.mean(dim=(-3, -2, -1), keepdim=True)
std_t = denoised.std(dim=(-3, -2, -1), keepdim=True)
ref_img_shift = ref_latent * std_t + avg_t
d = sampling.to_d(x, sigma, ref_img_shift)
d = to_d(x, sigma, ref_img_shift)
dt = (sigma_next - sigma) * factor
return x + d * dt
return (d * dt).add_(x)
@staticmethod
def guidance_linear(x: Tensor, ref_latent: Tensor, factor: float = 0.2) -> Tensor:
avg_t = x.mean(dim=[1, 2, 3], keepdim=True)
std_t = x.std(dim=[1, 2, 3], keepdim=True)
ref_img_shift = ref_latent * std_t + avg_t
return (1.0 - factor) * x + factor * ref_img_shift
def guidance_linear(
x: Tensor,
ref_latent: Tensor,
factor: float = 0.2,
*,
blend=torch.lerp,
) -> Tensor:
avg_t = x.mean(dim=(-3, -2, -1), keepdim=True)
std_t = x.std(dim=(-3, -2, -1), keepdim=True)
ref_img_shift = (ref_latent * std_t).add_(avg_t)
return blend(x, ref_img_shift, factor)
class SonarWithGuidance(SonarBase, SonarGuidanceMixin):
@@ -210,9 +409,9 @@ class SonarSampler(SonarWithGuidance):
def __init__(
self,
model,
sigmas,
s_in,
extra_args,
sigmas: Tensor,
s_in: Tensor,
extra_args: dict[str, Any],
*args: list[Any],
**kwargs: dict[str, Any],
):
@@ -222,62 +421,49 @@ class SonarSampler(SonarWithGuidance):
self.s_in = s_in
self.extra_args = extra_args
def call_model(
self,
x: Tensor,
sigma: Tensor,
*args: list[Any],
s_in=None,
extra_args=None,
) -> Tensor:
if s_in is None:
s_in = self.s_in
extra_args = (
self.extra_args if extra_args is None else self.extra_args | extra_args
)
return self.model(x, sigma * s_in, *args, **extra_args)
class SonarEuler(SonarSampler):
def __init__(
self,
s_churn: float = 0.0,
s_tmin: float = 0.0,
s_tmax: float = float("inf"),
s_noise: float = 1.0,
*args: list[Any],
**kwargs: dict[str, Any],
):
super().__init__(*args, **kwargs)
self.s_churn = s_churn
self.s_tmin = s_tmin
self.s_tmax = s_tmax
self.s_noise = s_noise
def step(
self,
step_index: int,
sample: torch.FloatTensor,
):
self.init_hist_d(sample)
def step(self, step_index: int, sample: torch.FloatTensor):
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
sigma, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1]
gamma = (
min(self.s_churn / (len(self.sigmas) - 1), 2**0.5 - 1)
if self.s_tmin <= sigma <= self.s_tmax
else 0.0
denoised = self.call_model(sample, sigma)
result_sample = self.momentum_step(
step_index,
sample,
denoised,
sigma,
sigma_next,
)
sigma_hat = sigma * (gamma + 1)
if gamma > 0:
noise = (
self.noise_sampler(sigma, sigma_to)
if self.noise_sampler
else torch.randn_like(sample)
)
eps = noise * self.s_noise
sample = sample + eps * (sigma_hat**2 - sigma**2) ** 0.5 # noqa: PLR6104
denoised = self.model(sample, sigma_hat * self.s_in, **self.extra_args)
derivative = sampling.to_d(sample, sigma, denoised)
dt = self.sigmas[step_index + 1] - sigma_hat
result_sample = self.momentum_step(sample, derivative, dt)
if self.sigmas[step_index + 1] > 0:
if sigma_next > 0:
result_sample = self.guidance_step(step_index, result_sample, denoised)
return (
result_sample,
sigma,
sigma_hat,
sigma,
denoised,
)
@@ -286,26 +472,18 @@ class SonarEuler(SonarSampler):
def sampler(
cls,
model,
x,
sigmas,
extra_args=None,
x: Tensor,
sigmas: Tensor,
extra_args: dict | None = None,
callback=None,
disable=None,
disable: bool | None = None, # noqa: FBT001
noise_sampler: Callable | None = None,
sonar_config=None,
s_churn=0.0,
s_tmin=0.0,
s_tmax=float("inf"),
s_noise=1.0,
):
if sonar_config is None:
sonar_config = SonarConfig()
s_in = x.new_ones([x.shape[0]])
sonar_config: SonarConfig | None = None,
sonar_params: dict | None = None,
) -> Tensor:
sonar_config = cls.get_config(sonar_config, sonar_params)
s_in = x.new_ones((x.shape[0],))
sonar = cls(
s_churn,
s_tmin,
s_tmax,
s_noise,
model,
sigmas,
s_in,
@@ -320,7 +498,7 @@ class SonarEuler(SonarSampler):
)
for i in trange(len(sigmas) - 1, disable=disable):
x, _sigma, sigma_hat, denoised = sonar.step(
x, sigma, sigma_hat, denoised = sonar.step(
i,
x,
)
@@ -329,7 +507,7 @@ class SonarEuler(SonarSampler):
{
"x": x,
"i": i,
"sigma": sigmas[i],
"sigma": sigma,
"sigma_hat": sigma_hat,
"denoised": denoised,
},
@@ -354,31 +532,32 @@ class SonarEulerAncestral(SonarSampler):
step_index: int,
sample: torch.FloatTensor,
):
self.init_hist_d(sample)
sigma_from, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1]
sigma_down, sigma_up = sampling.get_ancestral_step(
sigma_from,
sigma_to,
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
sigma_down, sigma_up = get_ancestral_step(
sigma,
sigma_next,
eta=self.eta,
)
denoised = self.model(sample, sigma_from * self.s_in, **self.extra_args)
derivative = sampling.to_d(sample, sigma_from, denoised)
dt = sigma_down - sigma_from
result_sample = self.momentum_step(sample, derivative, dt)
if sigma_to > 0:
denoised = self.call_model(sample, sigma)
result_sample = self.momentum_step(
step_index,
sample,
denoised,
sigma,
sigma_down,
)
if sigma_next > 0:
result_sample = self.guidance_step(step_index, result_sample, denoised)
result_sample = ( # noqa: PLR6104
result_sample
+ self.noise_sampler(sigma_from, sigma_to) * self.s_noise * sigma_up
+ self.noise_sampler(sigma, sigma_next) * (self.s_noise * sigma_up)
)
return (
result_sample,
sigma_from,
sigma_from,
sigma,
sigma,
denoised,
)
@@ -392,14 +571,14 @@ class SonarEulerAncestral(SonarSampler):
extra_args=None,
callback=None,
disable=None,
sonar_config=None,
sonar_config: SonarConfig | None = None,
sonar_params: dict | None = None,
eta=1.0,
s_noise=1.0,
noise_sampler: Callable | None = None,
):
if sonar_config is None:
sonar_config = SonarConfig()
s_in = x.new_ones([x.shape[0]])
sonar_config = cls.get_config(sonar_config, sonar_params)
s_in = x.new_ones((x.shape[0],))
sonar = cls(
eta,
s_noise,
@@ -449,104 +628,134 @@ class SonarDPMPPSDE(SonarSampler):
self.s_noise = s_noise
@staticmethod
def sigma_fn(t) -> float:
def sigma_fn(t: Tensor) -> float:
return t.neg().exp()
@staticmethod
def t_fn(sigma) -> float:
return sigma.log.neg()
def t_fn(sigma: Tensor) -> float:
return sigma.log().neg()
# DPM++ solver algorithm copied from ComfyUI source.
def momentum_step( # noqa: PLR0914
self,
step_index,
step_index: int,
x: Tensor,
denoised: Tensor,
sigma_from,
sigma_to,
sigma_down,
):
if sigma_to == 0:
derivative = sampling.to_d(x, sigma_from, denoised)
dt = sigma_down - sigma_from
return super().momentum_step(x, derivative, dt)
sigma: Tensor,
sigma_next: Tensor,
sigma_down: Tensor,
) -> Tensor:
if sigma_next == 0:
return super().momentum_step(step_index, x, denoised, sigma, sigma_down)
def sigma_fn(t):
return t.neg().exp()
def t_fn(sigma):
return sigma.log().neg()
hd = self.history_d
p = (1.0 - self.cfg.momentum) * self.cfg.direction
cfg = self.cfg
# Halve the momentum proportion if there's history since we will use it twice.
adjusted_momentum = (
cfg.momentum + (1 - cfg.momentum) / 2
if self.history_d is not None
else cfg.momentum
)
r = 1 / 2
# DPM-Solver++
t, t_next = t_fn(sigma_from), t_fn(sigma_to)
t, t_next = self.t_fn(sigma), self.t_fn(sigma_next)
h = t_next - t
s = t + h * r
fac = 1 / (2 * r)
# Step 1
sd, su = sampling.get_ancestral_step(sigma_fn(t), sigma_fn(s), self.eta)
s_ = t_fn(sd)
diff_2 = (t - s_).expm1() * denoised
momentum_d = (1.0 - p) * diff_2 + p * hd
self.update_hist(momentum_d)
hd = self.history_d
x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - momentum_d
x_2 += self.noise_sampler(sigma_fn(t), sigma_fn(s)) * self.s_noise * su
denoised_2 = self.model(x_2, sigma_fn(s) * self.s_in, **self.extra_args)
# Step 2
sd, su = sampling.get_ancestral_step(
sigma_fn(t),
sigma_fn(t_next),
s_t, s_s = self.sigma_fn(t), self.sigma_fn(s)
sd, su = get_ancestral_step(
s_t,
s_s,
self.eta,
)
t_next_ = t_fn(sd)
denoised_d = (1 - fac) * denoised + fac * denoised_2
diff_1 = (t - t_next_).expm1() * denoised_d
momentum_d = (1.0 - p) * diff_1 + p * hd
self.update_hist(momentum_d)
x = (sigma_fn(t_next_) / sigma_fn(t)) * x - momentum_d
s_ = self.t_fn(sd)
momentum_denoised = self.get_momentum_denoised(
x,
denoised,
sigma,
step=step_index,
)
diff_2 = (t - s_).expm1() * momentum_denoised
momentum_d = self.get_momentum_d(
x,
momentum_denoised,
sigma,
step=step_index,
momentum=adjusted_momentum,
d=diff_2,
)
x_2 = ((self.sigma_fn(s_) / s_t) * x).sub_(momentum_d)
x_2 += self.noise_sampler(s_t, s_s).mul_(
self.s_noise * su,
)
sigma_2 = s_s
denoised_2 = self.call_model(x_2, sigma_2)
momentum_denoised_2 = self.get_momentum_denoised(
x,
denoised_2,
sigma_2,
step=step_index,
)
# Step 2
s_t_next = self.sigma_fn(t_next)
sd, su = get_ancestral_step(
s_t,
s_t_next,
self.eta,
)
t_down = self.t_fn(sd)
denoised_d = (1 - fac) * momentum_denoised + fac * momentum_denoised_2
diff_1 = (t - t_down).expm1() * denoised_d
momentum_d = self.get_momentum_d(
x,
momentum_denoised_2,
sigma_2,
step=step_index,
momentum=adjusted_momentum,
d=diff_1,
)
x = ((self.sigma_fn(t_down) / s_t) * x).sub_(momentum_d)
x = self.guidance_step(step_index, x, denoised_d)
return x + self.noise_sampler(sigma_fn(t), sigma_fn(t_next)) * self.s_noise * su
x += self.noise_sampler(s_t, s_t_next).mul_(
self.s_noise * su,
)
return x
def step(
self,
step_index: int,
sample: torch.FloatTensor,
):
) -> Tensor:
def sigma_fn(t):
return t.neg().exp()
def t_fn(sigma):
return sigma.log().neg()
self.init_hist_d(sample)
sigma_from, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1]
sigma_down, _sigma_up = sampling.get_ancestral_step(
sigma_from,
sigma_to,
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
sigma_down, _sigma_up = get_ancestral_step(
sigma,
sigma_next,
eta=self.eta,
)
denoised = self.model(sample, sigma_from * self.s_in, **self.extra_args)
denoised = self.call_model(sample, sigma)
result_sample = self.momentum_step(
step_index,
sample,
denoised,
sigma_from,
sigma_to,
sigma,
sigma_next,
sigma_down,
)
return (
result_sample,
sigma_from,
sigma_from,
sigma,
sigma,
denoised,
)
@@ -555,19 +764,19 @@ class SonarDPMPPSDE(SonarSampler):
def sampler(
cls,
model,
x,
sigmas,
extra_args=None,
x: Tensor,
sigmas: Tensor,
extra_args: dict | None = None,
callback=None,
disable=None,
sonar_config=None,
disable: bool | None = None, # noqa: FBT001
sonar_config: SonarConfig | None = None,
sonar_params: dict | None = None,
eta=1.0,
s_noise=1.0,
noise_sampler=None,
):
if sonar_config is None:
sonar_config = SonarConfig()
s_in = x.new_ones([x.shape[0]])
) -> Tensor:
sonar_config = cls.get_config(sonar_config, sonar_params)
s_in = x.new_ones((x.shape[0],))
sonar = cls(
eta,
s_noise,
@@ -602,7 +811,7 @@ class SonarDPMPPSDE(SonarSampler):
return x
def add_samplers():
def add_samplers() -> None:
extra_samplers = {
"sonar_euler": SonarEuler.sampler,
"sonar_euler_ancestral": SonarEulerAncestral.sampler,
+196
View File
@@ -0,0 +1,196 @@
from __future__ import annotations
import math
import torch
from comfy.model_management import device_supports_non_blocking
from comfy.utils import common_upscale
from .external import MODULES as EXT
BLENDING_MODES = {"lerp": torch.lerp}
UPSCALE_METHODS = (
"bilinear",
"nearest-exact",
"nearest",
"area",
"bicubic",
"bislerp",
)
def scale_samples(
samples: torch.Tensor,
width: int,
height: int,
*,
mode: str = "bicubic",
) -> torch.Tensor:
return common_upscale(samples, width, height, mode, None)
def init_integrations(integrations) -> None:
global scale_samples, BLENDING_MODES, UPSCALE_METHODS # noqa: PLW0603
bleh = integrations.bleh
if bleh is None:
return
bleh_latentutils = bleh.py.latent_utils
BLENDING_MODES = bleh_latentutils.BLENDING_MODES
UPSCALE_METHODS = bleh_latentutils.UPSCALE_METHODS
scale_samples = bleh_latentutils.scale_samples
EXT.register_init_handler(init_integrations)
def scale_noise(
noise: torch.Tensor,
factor: float = 1.0,
*,
normalized: bool = True,
threshold_std_devs: float = 2.5,
normalize_dims: tuple | None = None,
) -> torch.Tensor:
numel = noise.numel()
if not normalized or numel == 0:
return noise.mul_(factor) if factor != 1 else noise
if normalize_dims is not None:
std = noise.std(dim=normalize_dims, keepdim=True)
noise = noise / std # noqa: PLR6104
return noise.sub_(noise.mean(dim=normalize_dims, keepdim=True)).mul_(factor)
mean, std = noise.mean().item(), noise.std().item()
threshold = threshold_std_devs / math.sqrt(numel)
if abs(mean) > threshold:
noise -= mean
if abs(1.0 - std) > threshold:
noise /= std
return noise.mul_(factor) if factor != 1 else noise
CAN_NONBLOCK = {}
def tensor_to(
tensor: torch.Tensor,
dest: torch.Tensor | torch.Device | str,
) -> torch.Tensor:
device = dest.device if isinstance(dest, torch.Tensor) else dest
non_blocking = CAN_NONBLOCK.get(device)
if non_blocking is None:
non_blocking = device_supports_non_blocking(device)
CAN_NONBLOCK[device] = non_blocking
return tensor.to(dest, non_blocking=non_blocking)
def quantile_normalize(
noise: torch.Tensor,
*,
quantile: float = 0.75,
dim: int | None = 1,
flatten: bool = True,
nq_fac: float = 1.0,
pow_fac: float = 0.5,
) -> torch.Tensor:
if quantile is None or quantile <= 0 or quantile >= 1:
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 qdim in {0, 1}:
flatdim = qdim + 1
elif qdim in {2, 3}:
noise = noise.movedim(qdim, 1)
tempshape = noise.shape
flatdim = 2
else:
raise ValueError(
"Cannot handling quantile normalization flattening dims > 3",
)
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)
noise = noise.clamp(-nq, nq)
noise = torch.copysign(
torch.pow(torch.abs(noise), pow_fac),
noise,
)
if flatdim is not None and qdim in {2, 3}:
return (
noise.reshape(tempshape).movedim(1, qdim).reshape(orig_shape).contiguous()
)
return noise
def adjust_slice(s: slice, size: int, offset: int) -> slice:
if offset == 0:
return s
# Input slice must have positive start/stop and be in bounds for the object that will be sliced here.
start = s.start if s.start is not None else 0
stop = s.stop if s.stop is not None else size
if offset < 0:
adj = min(start, abs(offset))
return slice(start - adj, stop - adj)
adj = min(size - stop, offset)
return slice(start + adj, stop + adj)
def crop_samples(
tensor: torch.Tensor,
width: int,
height: int,
*,
mode="center",
offset_width: int = 0,
offset_height: int = 0,
):
if tensor.ndim < 3:
raise ValueError("Can only handle >= 3 dimensional tensors")
th, tw = tensor.shape[-2:]
if (tw, th) == (width, height):
return tensor
if tw < width or th < height:
raise ValueError("Can't crop sample smaller than requested width or height")
if mode == "center":
hmode = wmode = "center"
else:
hmode, wmode, *splitextra = mode.split("_")
if splitextra:
raise ValueError("Bad composite mode")
if hmode == "top":
hslice = slice(0, height)
elif hmode == "center":
hoffs = (th - height) // 2
hslice = slice(hoffs, hoffs + height)
elif hmode == "bottom":
hslice = slice(th - height, th)
else:
raise ValueError("Bad height mode in composite mode")
if wmode == "left":
wslice = slice(0, width)
elif wmode == "center":
woffs = (tw - width) // 2
wslice = slice(woffs, woffs + width)
elif wmode == "right":
wslice = slice(tw - width, tw)
else:
raise ValueError("Bad width mode in composite mode")
wslice = adjust_slice(wslice, tw, offset_width)
hslice = adjust_slice(hslice, th, offset_height)
return tensor[..., hslice, wslice]
+2
View File
@@ -15,6 +15,8 @@ ignore = [
"D102",
"D103",
"D104",
"D105",
"D106",
"D107",
"D211",
"D213",