Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a4fed311a8 | ||
|
|
2988afa34a | ||
|
|
68fc7418d1 | ||
|
|
6d15c0bbca | ||
|
|
f7cbbfcbda |
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
+521
-233
@@ -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
File diff suppressed because it is too large
Load Diff
+4
-3
@@ -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
@@ -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
@@ -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]
|
||||
Reference in New Issue
Block a user