Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e36623a5f1 | ||
|
|
ca3ee58750 | ||
|
|
a31eb6940b | ||
|
|
3222b02318 | ||
|
|
dcfea85e9c | ||
|
|
3ba9f2e3d1 | ||
|
|
b5be44720c | ||
|
|
30b37e98c2 | ||
|
|
a951ad7392 | ||
|
|
27126d9f93 | ||
|
|
7365a9f30b |
@@ -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...
|
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
|
## 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.
|
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.
|
||||||
@@ -112,6 +153,7 @@ My version was initially based on this Sonar sampler implementation for Diffuser
|
|||||||
* 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.
|
* 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.
|
* 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
|
* 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
|
## Errata
|
||||||
|
|
||||||
|
|||||||
@@ -2,6 +2,22 @@
|
|||||||
|
|
||||||
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
|
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
|
||||||
|
|
||||||
|
## 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
|
## 20241129
|
||||||
|
|
||||||
*Note*: Contains some potentially workflow-breaking changes.
|
*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`
|
### `NoisyLatentLike`
|
||||||
|
|
||||||
This node takes a reference latent and generates noise of the same shape. The one required input is `latent`.
|
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`
|
### `SonarModulatedNoise`
|
||||||
|
|
||||||
Experimental noise modulation based on code stolen from
|
Experimental noise modulation based on code stolen from
|
||||||
|
|||||||
+117
-14
@@ -1,22 +1,125 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import contextlib
|
import contextlib
|
||||||
import importlib
|
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):
|
class Integrations:
|
||||||
import custom_nodes.ComfyUI_restart_sampling as rs
|
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",)
|
__all__ = ("MODULES",)
|
||||||
|
|||||||
+7
-13
@@ -2,7 +2,8 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from .external import MODULES as EXTERNAL_MODULES
|
from . import utils
|
||||||
|
from .external import IntegratedNode
|
||||||
from .powernoise import PowerFilter
|
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)
|
return x_filt.to(x.dtype, non_blocking=True)
|
||||||
|
|
||||||
|
|
||||||
BLEND_OPS = (
|
class FreeUExtremeConfigNode(metaclass=IntegratedNode):
|
||||||
{"lerp": torch.lerp}
|
|
||||||
if "bleh" not in EXTERNAL_MODULES
|
|
||||||
else EXTERNAL_MODULES["bleh"].py.latent_utils.BLENDING_MODES
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class FreeUExtremeConfigNode:
|
|
||||||
DESCRIPTION = "Allows setting configuration for FreeU Extreme."
|
DESCRIPTION = "Allows setting configuration for FreeU Extreme."
|
||||||
RETURN_TYPES = ("FRUX_CONFIG",)
|
RETURN_TYPES = ("FRUX_CONFIG",)
|
||||||
FUNCTION = "go"
|
FUNCTION = "go"
|
||||||
@@ -150,7 +144,7 @@ class FreeUExtremeConfigNode:
|
|||||||
},
|
},
|
||||||
),
|
),
|
||||||
"blend_mode": (
|
"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",
|
"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] = (
|
x[:, slice_offs : slice_offs + slice_size] = (
|
||||||
xslice
|
xslice
|
||||||
if self.blend == 1.0
|
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],
|
x[:, slice_offs : slice_offs + slice_size],
|
||||||
xslice,
|
xslice,
|
||||||
self.blend,
|
self.blend,
|
||||||
@@ -331,12 +325,12 @@ class FreeUExtremeConfig:
|
|||||||
def clone(self):
|
def clone(self):
|
||||||
return self.__class__(**{k: getattr(self, k) for k in self._keys})
|
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}
|
meh = {k: getattr(self, k) for k in self._keys}
|
||||||
return f"<FRUXConfig: {meh}>"
|
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."
|
DESCRIPTION = "Main FreeU Extreme node. Allows patching a model with the FreeU (V2) effect with more control."
|
||||||
RETURN_TYPES = ("MODEL",)
|
RETURN_TYPES = ("MODEL",)
|
||||||
FUNCTION = "go"
|
FUNCTION = "go"
|
||||||
|
|||||||
+733
-357
File diff suppressed because it is too large
Load Diff
+519
-233
@@ -6,18 +6,30 @@ from typing import Callable
|
|||||||
|
|
||||||
import comfy
|
import comfy
|
||||||
import torch
|
import torch
|
||||||
|
import yaml
|
||||||
from comfy.k_diffusion import sampling
|
from comfy.k_diffusion import sampling
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from . import external
|
from . import external, utils
|
||||||
from .noise_generation import *
|
from .noise_generation import *
|
||||||
from .sonar import SonarGuidanceMixin
|
from .sonar import SonarGuidanceMixin
|
||||||
|
from .utils import crop_samples, scale_noise
|
||||||
|
|
||||||
# ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311
|
# ruff: noqa: D412, D413, D417, D212, D407, ANN002, ANN003, FBT001, FBT002, S311
|
||||||
|
|
||||||
|
|
||||||
class CustomNoiseItemBase(abc.ABC):
|
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.factor = factor
|
||||||
self.keys = set(kwargs.keys())
|
self.keys = set(kwargs.keys())
|
||||||
for k, v in kwargs.items():
|
for k, v in kwargs.items():
|
||||||
@@ -46,6 +58,7 @@ class CustomNoiseItemBase(abc.ABC):
|
|||||||
seed=None,
|
seed=None,
|
||||||
cpu=True,
|
cpu=True,
|
||||||
normalized=True,
|
normalized=True,
|
||||||
|
**kwargs,
|
||||||
):
|
):
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@@ -65,16 +78,25 @@ class CustomNoiseItem(CustomNoiseItemBase):
|
|||||||
seed=None,
|
seed=None,
|
||||||
cpu=True,
|
cpu=True,
|
||||||
normalized=True,
|
normalized=True,
|
||||||
|
**kwargs,
|
||||||
):
|
):
|
||||||
|
ns_kwargs = getattr(self, "ns_kwargs", {}).copy()
|
||||||
|
# print("NS KWARGS", ns_kwargs)
|
||||||
|
|
||||||
return get_noise_sampler(
|
return get_noise_sampler(
|
||||||
self.noise_type,
|
self.noise_type,
|
||||||
x,
|
x,
|
||||||
sigma_min,
|
sigma_min,
|
||||||
sigma_max,
|
sigma_max,
|
||||||
seed=seed,
|
seed=ns_kwargs.pop("seed", seed),
|
||||||
cpu=cpu,
|
cpu=ns_kwargs.pop("cpu", cpu),
|
||||||
factor=self.factor,
|
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,
|
make_noise_sampler: Callable | None = None,
|
||||||
normalized=False,
|
normalized=False,
|
||||||
factor: float = 1.0,
|
factor: float = 1.0,
|
||||||
|
**kwargs,
|
||||||
):
|
):
|
||||||
self.factor = factor
|
self.factor = factor
|
||||||
self.normalized = normalized
|
self.normalized = normalized
|
||||||
@@ -164,16 +187,18 @@ class NoiseSampler:
|
|||||||
try:
|
try:
|
||||||
self.noise_sampler = make_noise_sampler(
|
self.noise_sampler = make_noise_sampler(
|
||||||
x,
|
x,
|
||||||
transform(torch.as_tensor(sigma_min))
|
sigma_min=transform(torch.as_tensor(sigma_min))
|
||||||
if sigma_min is not None
|
if sigma_min is not None
|
||||||
else None,
|
else None,
|
||||||
transform(torch.as_tensor(sigma_max))
|
sigma_max=transform(torch.as_tensor(sigma_max))
|
||||||
if sigma_max is not None
|
if sigma_max is not None
|
||||||
else None,
|
else None,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
cpu=cpu,
|
cpu=cpu,
|
||||||
|
**kwargs,
|
||||||
)
|
)
|
||||||
except TypeError as _exc:
|
except TypeError as _exc:
|
||||||
|
print("GOT EXC", _exc)
|
||||||
self.noise_sampler = make_noise_sampler(x)
|
self.noise_sampler = make_noise_sampler(x)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -216,7 +241,7 @@ class AdvancedNoiseBase(CustomNoiseItemBase):
|
|||||||
v = getattr(self, k, None)
|
v = getattr(self, k, None)
|
||||||
if v is not None:
|
if v is not None:
|
||||||
noise_sampler_kwargs[k] = v
|
noise_sampler_kwargs[k] = v
|
||||||
self.sampler_factory = NoiseSampler.simple(
|
self.sampler_factory = NoiseSampler.wrap(
|
||||||
partial(self.ns_factory, **noise_sampler_kwargs),
|
partial(self.ns_factory, **noise_sampler_kwargs),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -229,9 +254,9 @@ class AdvancedPyramidNoise(AdvancedNoiseBase):
|
|||||||
ns_factory_arg_keys = ("discount", "iterations", "upscale_mode")
|
ns_factory_arg_keys = ("discount", "iterations", "upscale_mode")
|
||||||
|
|
||||||
pyramid_variants_map = { # noqa: RUF012
|
pyramid_variants_map = { # noqa: RUF012
|
||||||
"pyramid": pyramid_noise_like,
|
"pyramid": PyramidNoiseGenerator,
|
||||||
"pyramid_old": pyramid_old_noise_like,
|
"pyramid_old": PyramidOldNoiseGenerator,
|
||||||
"highres_pyramid": highres_pyramid_noise_like,
|
"highres_pyramid": HighresPyramidNoiseGenerator,
|
||||||
}
|
}
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -244,15 +269,33 @@ class Advanced1fNoise(AdvancedNoiseBase):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def ns_factory(self):
|
def ns_factory(self):
|
||||||
return onef_noise_like
|
return OneFNoiseGenerator
|
||||||
|
|
||||||
|
|
||||||
class AdvancedPowerLawNoise(AdvancedNoiseBase):
|
class AdvancedPowerLawNoise(AdvancedNoiseBase):
|
||||||
ns_factory_arg_keys = ("alpha", "div_max_dims", "use_sign")
|
ns_factory_arg_keys = ("alpha", "div_max_dims", "use_sign")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
@property
|
@property
|
||||||
def ns_factory(self):
|
def ns_factory(cls):
|
||||||
return powerlaw_noise_like
|
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):
|
class CompositeNoise(CustomNoiseItemBase):
|
||||||
@@ -446,6 +489,10 @@ class ScheduledNoise(CustomNoiseItemBase):
|
|||||||
return torch.zeros_like(x)
|
return torch.zeros_like(x)
|
||||||
|
|
||||||
def noise_sampler(s, sn):
|
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)
|
noise = (ns if end_sigma <= s <= start_sigma else nsa)(s, sn)
|
||||||
return scale_noise(noise, factor, normalized=normalize)
|
return scale_noise(noise, factor, normalized=normalize)
|
||||||
|
|
||||||
@@ -997,271 +1044,508 @@ class BlendedNoise(CustomNoiseItemBase):
|
|||||||
return noise_sampler
|
return noise_sampler
|
||||||
|
|
||||||
|
|
||||||
if "bleh" in external.MODULES:
|
class ResizedNoise(CustomNoiseItemBase):
|
||||||
bleh = external.MODULES["bleh"]
|
def __init__(
|
||||||
BLU = bleh.py.latent_utils
|
self,
|
||||||
BOPS = bleh.py.nodes.ops
|
factor,
|
||||||
|
*,
|
||||||
class BlendFilterNoise(CustomNoiseItemBase):
|
width,
|
||||||
def __init__(
|
height,
|
||||||
self,
|
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,
|
factor,
|
||||||
*,
|
width=width,
|
||||||
noise,
|
height=height,
|
||||||
blend_mode,
|
downscale_strategy=downscale_strategy,
|
||||||
ffilter,
|
initial_reference=initial_reference,
|
||||||
ffilter_scale,
|
crop_offset_horizontal=crop_offset_horizontal,
|
||||||
ffilter_strength,
|
crop_offset_vertical=crop_offset_vertical,
|
||||||
ffilter_threshold,
|
crop_mode=crop_mode,
|
||||||
enhance_mode,
|
upscale_mode=upscale_mode,
|
||||||
enhance_strength,
|
downscale_mode=downscale_mode,
|
||||||
affect,
|
custom_noise=custom_noise.clone(),
|
||||||
normalize_result,
|
normalize=normalize,
|
||||||
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):
|
def clone_key(self, k):
|
||||||
if k == "noise":
|
if k == "custom_noise":
|
||||||
return self.noise.clone()
|
return self.custom_noise.clone()
|
||||||
return super().clone_key(k)
|
return super().clone_key(k)
|
||||||
|
|
||||||
def apply_effects(self, noise, sigma):
|
def make_noise_sampler(self, x, *args, normalized=True, **kwargs):
|
||||||
if self.ffilter:
|
if x.ndim < 3:
|
||||||
noise = BLU.ffilter(
|
raise ValueError("ResizedNoise can only handle 3+ dimensional latents")
|
||||||
noise,
|
factor = self.factor
|
||||||
self.ffilter_threshold,
|
normalize = self.get_normalize("normalize", normalized)
|
||||||
self.ffilter_scale,
|
xh, xw = x.shape[-2:]
|
||||||
self.ffilter,
|
nh, nw = self.height // 8, self.width // 8
|
||||||
self.ffilter_strength,
|
offsh, offsw = self.crop_offset_vertical // 8, self.crop_offset_horizontal // 8
|
||||||
)
|
if xh == nh and xw == nw:
|
||||||
if self.enhance_mode != "none" and self.enhance_strength != 0:
|
ns = self.custom_noise.make_noise_sampler(
|
||||||
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(
|
|
||||||
x,
|
x,
|
||||||
*args,
|
*args,
|
||||||
normalized=False,
|
normalized=normalize,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
def noise_sampler(s, sn):
|
def noise_sampler(*args, **kwargs):
|
||||||
noise = internal_ns(s, sn)
|
return ns(*args, **kwargs).mul_(factor)
|
||||||
if len(rulegroup.rules):
|
|
||||||
state = {
|
|
||||||
BOPS.CondType.TYPE: BOPS.PatchType.LATENT,
|
|
||||||
BOPS.CondType.PERCENT: 0.0,
|
|
||||||
BOPS.CondType.BLOCK: -1,
|
|
||||||
BOPS.CondType.STAGE: -1,
|
|
||||||
"sigma": None if s is None else s,
|
|
||||||
"h": noise,
|
|
||||||
"hsp": x.detach().clone(),
|
|
||||||
"target": "h",
|
|
||||||
}
|
|
||||||
noise = rulegroup.eval(state, toplevel=True)["h"]
|
|
||||||
return scale_noise(noise, factor, normalized=normalize)
|
|
||||||
|
|
||||||
return noise_sampler
|
return noise_sampler
|
||||||
|
|
||||||
|
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] = {
|
NOISE_SAMPLERS: dict[NoiseType, Callable] = {
|
||||||
NoiseType.BROWNIAN: NoiseSampler.wrap(sampling.BrownianTreeNoiseSampler),
|
NoiseType.BROWNIAN: NoiseSampler.wrap(BrownianNoiseGenerator),
|
||||||
NoiseType.GAUSSIAN: NoiseSampler.simple(torch.randn_like),
|
NoiseType.DISTRO: NoiseSampler.wrap(DistroNoiseGenerator),
|
||||||
NoiseType.UNIFORM: NoiseSampler.simple(uniform_noise_like),
|
NoiseType.GAUSSIAN: NoiseSampler.wrap(GaussianNoiseGenerator),
|
||||||
NoiseType.PERLIN: NoiseSampler.simple(rand_perlin_like),
|
NoiseType.UNIFORM: NoiseSampler.wrap(UniformNoiseGenerator),
|
||||||
NoiseType.STUDENTT: NoiseSampler.simple(studentt_noise_like),
|
NoiseType.PERLIN: NoiseSampler.wrap(PerlinOldNoiseGenerator),
|
||||||
NoiseType.ONEF_PINKISH: NoiseSampler.simple(partial(onef_noise_like, alpha=-0.5)),
|
NoiseType.STUDENTT: NoiseSampler.wrap(StudentTNoiseGenerator),
|
||||||
NoiseType.ONEF_GREENISH: NoiseSampler.simple(partial(onef_noise_like, alpha=0.5)),
|
NoiseType.ONEF_PINKISH: NoiseSampler.wrap(partial(OneFNoiseGenerator, alpha=-0.5)),
|
||||||
NoiseType.ONEF_PINKISHGREENISH: NoiseSampler.simple(
|
NoiseType.ONEF_GREENISH: NoiseSampler.wrap(partial(OneFNoiseGenerator, alpha=0.5)),
|
||||||
lambda x: onef_noise_like(x, alpha=0.5)
|
NoiseType.ONEF_PINKISHGREENISH: NoiseSampler.wrap(
|
||||||
.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(
|
|
||||||
partial(
|
partial(
|
||||||
powerlaw_noise_like,
|
MixedNoiseGenerator,
|
||||||
|
name="onef_pinkishgreenish",
|
||||||
|
noise_mix=(
|
||||||
|
partial(OneFNoiseGenerator, alpha=0.5),
|
||||||
|
partial(OneFNoiseGenerator, alpha=-0.5),
|
||||||
|
),
|
||||||
|
output_fun=lambda t: t.mul_(0.5),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
NoiseType.ONEF_PINKISH_MIX: NoiseSampler.wrap(
|
||||||
|
partial(
|
||||||
|
MixedNoiseGenerator,
|
||||||
|
name="onef_pinkish_mix",
|
||||||
|
noise_mix=(
|
||||||
|
(partial(OneFNoiseGenerator, alpha=0.5), lambda t: t.mul_(-1.0)),
|
||||||
|
partial(OneFNoiseGenerator, alpha=0.5),
|
||||||
|
),
|
||||||
|
output_fun=lambda t: t.mul_(0.5),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
NoiseType.ONEF_GREENISH_MIX: NoiseSampler.wrap(
|
||||||
|
partial(
|
||||||
|
MixedNoiseGenerator,
|
||||||
|
name="onef_greenish_mix",
|
||||||
|
noise_mix=(
|
||||||
|
(partial(OneFNoiseGenerator, alpha=0.5), lambda t: t.mul_(-1.0)),
|
||||||
|
partial(OneFNoiseGenerator, alpha=0.5),
|
||||||
|
),
|
||||||
|
output_fun=lambda t: t.mul_(0.5),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
NoiseType.WHITE: NoiseSampler.wrap(
|
||||||
|
partial(
|
||||||
|
PowerLawNoiseGenerator,
|
||||||
alpha=0.0,
|
alpha=0.0,
|
||||||
use_sign=True,
|
use_sign=True,
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
NoiseType.GREY: NoiseSampler.simple(
|
NoiseType.GREY: NoiseSampler.wrap(
|
||||||
partial(
|
partial(
|
||||||
powerlaw_noise_like,
|
PowerLawNoiseGenerator,
|
||||||
alpha=0.0,
|
alpha=0.0,
|
||||||
use_sign=False,
|
use_sign=False,
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
NoiseType.VELVET: NoiseSampler.simple(
|
NoiseType.VELVET: NoiseSampler.wrap(
|
||||||
partial(
|
partial(
|
||||||
powerlaw_noise_like,
|
PowerLawNoiseGenerator,
|
||||||
alpha=1.0,
|
alpha=1.0,
|
||||||
use_sign=True,
|
use_sign=True,
|
||||||
div_max_dims=(-3, -2, -1),
|
div_max_dims=(-3, -2, -1),
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
NoiseType.VIOLET: NoiseSampler.simple(
|
NoiseType.VIOLET: NoiseSampler.wrap(
|
||||||
partial(
|
partial(
|
||||||
powerlaw_noise_like,
|
PowerLawNoiseGenerator,
|
||||||
alpha=0.5,
|
alpha=0.5,
|
||||||
use_sign=True,
|
use_sign=True,
|
||||||
div_max_dims=(-3, -2, -1),
|
div_max_dims=(-3, -2, -1),
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
NoiseType.PINK_OLD: NoiseSampler.simple(pink_noise_old_like),
|
NoiseType.WAVELET: NoiseSampler.wrap(WaveletNoiseGenerator),
|
||||||
NoiseType.HIGHRES_PYRAMID: NoiseSampler.simple(highres_pyramid_noise_like),
|
NoiseType.PINK_OLD: NoiseSampler.wrap(PinkOldNoiseGenerator),
|
||||||
NoiseType.PYRAMID: NoiseSampler.simple(pyramid_noise_like),
|
NoiseType.HIGHRES_PYRAMID: NoiseSampler.wrap(HighresPyramidNoiseGenerator),
|
||||||
NoiseType.RAINBOW_MILD: NoiseSampler.simple(
|
NoiseType.PYRAMID: NoiseSampler.wrap(PyramidNoiseGenerator),
|
||||||
lambda x: green_noise_like(x)
|
NoiseType.RAINBOW_MILD: NoiseSampler.wrap(
|
||||||
.mul_(0.55)
|
partial(
|
||||||
.add_(rand_perlin_like(x).mul_(0.7))
|
MixedNoiseGenerator,
|
||||||
.mul_(1.15),
|
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(
|
NoiseType.RAINBOW_INTENSE: NoiseSampler.wrap(
|
||||||
lambda x: green_noise_like(x)
|
partial(
|
||||||
.mul_(0.75)
|
MixedNoiseGenerator,
|
||||||
.add_(rand_perlin_like(x).mul_(0.5))
|
name="rainbow_intense",
|
||||||
.mul_(1.15),
|
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.LAPLACIAN: NoiseSampler.wrap(LaplacianNoiseGenerator),
|
||||||
NoiseType.POWER_OLD: NoiseSampler.simple(power_noise_old_like),
|
NoiseType.POWER_OLD: NoiseSampler.wrap(PowerOldNoiseGenerator),
|
||||||
NoiseType.GREEN_TEST: NoiseSampler.simple(green_noise_like),
|
NoiseType.GREEN_TEST: NoiseSampler.wrap(GreenTestNoiseGenerator),
|
||||||
NoiseType.PYRAMID_OLD: NoiseSampler.simple(pyramid_old_noise_like),
|
NoiseType.PYRAMID_OLD: NoiseSampler.wrap(PyramidOldNoiseGenerator),
|
||||||
NoiseType.PYRAMID_BISLERP: NoiseSampler.simple(
|
NoiseType.PYRAMID_BISLERP: NoiseSampler.wrap(
|
||||||
partial(pyramid_noise_like, upscale_mode="bislerp"),
|
partial(PyramidNoiseGenerator, upscale_mode="bislerp"),
|
||||||
),
|
),
|
||||||
NoiseType.HIGHRES_PYRAMID_BISLERP: NoiseSampler.simple(
|
NoiseType.HIGHRES_PYRAMID_BISLERP: NoiseSampler.wrap(
|
||||||
partial(highres_pyramid_noise_like, upscale_mode="bislerp"),
|
partial(HighresPyramidNoiseGenerator, upscale_mode="bislerp"),
|
||||||
),
|
),
|
||||||
NoiseType.PYRAMID_AREA: NoiseSampler.simple(
|
NoiseType.PYRAMID_AREA: NoiseSampler.wrap(
|
||||||
partial(pyramid_noise_like, upscale_mode="area"),
|
partial(PyramidNoiseGenerator, upscale_mode="area"),
|
||||||
),
|
),
|
||||||
NoiseType.HIGHRES_PYRAMID_AREA: NoiseSampler.simple(
|
NoiseType.HIGHRES_PYRAMID_AREA: NoiseSampler.wrap(
|
||||||
partial(highres_pyramid_noise_like, upscale_mode="area"),
|
partial(HighresPyramidNoiseGenerator, upscale_mode="area"),
|
||||||
),
|
),
|
||||||
NoiseType.PYRAMID_OLD_BISLERP: NoiseSampler.simple(
|
NoiseType.PYRAMID_OLD_BISLERP: NoiseSampler.wrap(
|
||||||
partial(pyramid_old_noise_like, upscale_mode="bislerp"),
|
partial(PyramidOldNoiseGenerator, upscale_mode="bislerp"),
|
||||||
),
|
),
|
||||||
NoiseType.PYRAMID_OLD_AREA: NoiseSampler.simple(
|
NoiseType.PYRAMID_OLD_AREA: NoiseSampler.wrap(
|
||||||
partial(pyramid_old_noise_like, upscale_mode="area"),
|
partial(PyramidOldNoiseGenerator, upscale_mode="area"),
|
||||||
),
|
),
|
||||||
NoiseType.PYRAMID_DISCOUNT5: NoiseSampler.simple(
|
NoiseType.PYRAMID_DISCOUNT5: NoiseSampler.wrap(
|
||||||
partial(pyramid_noise_like, discount=0.5),
|
partial(PyramidNoiseGenerator, discount=0.5),
|
||||||
),
|
),
|
||||||
NoiseType.PYRAMID_MIX: NoiseSampler.simple(
|
NoiseType.PYRAMID_MIX: NoiseSampler.wrap(
|
||||||
lambda x: pyramid_noise_like(x, discount=0.6)
|
partial(
|
||||||
.mul_(0.2)
|
MixedNoiseGenerator,
|
||||||
.add_(pyramid_noise_like(x, discount=0.6).mul_(-0.8)),
|
name="pyramid_mix",
|
||||||
|
noise_mix=(
|
||||||
|
(partial(PyramidNoiseGenerator, discount=0.6), lambda t: t.mul_(0.2)),
|
||||||
|
(partial(PyramidNoiseGenerator, discount=0.6), lambda t: t.mul_(-0.8)),
|
||||||
|
),
|
||||||
|
),
|
||||||
),
|
),
|
||||||
NoiseType.PYRAMID_MIX_AREA: NoiseSampler.simple(
|
NoiseType.PYRAMID_MIX_AREA: NoiseSampler.wrap(
|
||||||
lambda x: pyramid_noise_like(x, discount=0.5, upscale_mode="area")
|
partial(
|
||||||
.mul_(0.2)
|
MixedNoiseGenerator,
|
||||||
.add_(pyramid_noise_like(x, discount=0.5, upscale_mode="area").mul_(-0.8)),
|
name="pyramid_mix_area",
|
||||||
|
noise_mix=(
|
||||||
|
(
|
||||||
|
partial(PyramidNoiseGenerator, discount=0.5, upscale_mode="area"),
|
||||||
|
lambda t: t.mul_(0.2),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
partial(PyramidNoiseGenerator, discount=0.5, upscale_mode="area"),
|
||||||
|
lambda t: t.mul_(-0.8),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
),
|
||||||
),
|
),
|
||||||
NoiseType.PYRAMID_MIX_BISLERP: NoiseSampler.simple(
|
NoiseType.PYRAMID_MIX_BISLERP: NoiseSampler.wrap(
|
||||||
lambda x: pyramid_noise_like(x, discount=0.5, upscale_mode="bislerp")
|
partial(
|
||||||
.mul_(0.2)
|
MixedNoiseGenerator,
|
||||||
.add_(pyramid_noise_like(x, discount=0.5, upscale_mode="bislerp").mul_(-0.8)),
|
name="pyramid_mix_bislerp",
|
||||||
|
noise_mix=(
|
||||||
|
(
|
||||||
|
partial(
|
||||||
|
PyramidNoiseGenerator,
|
||||||
|
discount=0.5,
|
||||||
|
upscale_mode="bislerp",
|
||||||
|
),
|
||||||
|
lambda t: t.mul_(0.2),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
partial(
|
||||||
|
PyramidNoiseGenerator,
|
||||||
|
discount=0.5,
|
||||||
|
upscale_mode="bislerp",
|
||||||
|
),
|
||||||
|
lambda t: t.mul_(-0.8),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
),
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1275,6 +1559,7 @@ def get_noise_sampler(
|
|||||||
cpu: bool = True,
|
cpu: bool = True,
|
||||||
factor: float = 1.0,
|
factor: float = 1.0,
|
||||||
normalized=False,
|
normalized=False,
|
||||||
|
**kwargs,
|
||||||
) -> Callable:
|
) -> Callable:
|
||||||
if noise_type is None:
|
if noise_type is None:
|
||||||
noise_type = NoiseType.GAUSSIAN
|
noise_type = NoiseType.GAUSSIAN
|
||||||
@@ -1293,4 +1578,5 @@ def get_noise_sampler(
|
|||||||
cpu=cpu,
|
cpu=cpu,
|
||||||
factor=factor,
|
factor=factor,
|
||||||
normalized=normalized,
|
normalized=normalized,
|
||||||
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|||||||
+1234
-441
File diff suppressed because it is too large
Load Diff
+1
-1
@@ -23,7 +23,7 @@ from .nodes import (
|
|||||||
SonarNormalizeNoiseNodeMixin,
|
SonarNormalizeNoiseNodeMixin,
|
||||||
)
|
)
|
||||||
from .noise import CustomNoiseItemBase
|
from .noise import CustomNoiseItemBase
|
||||||
from .noise_generation import scale_noise
|
from .utils import scale_noise
|
||||||
|
|
||||||
# ruff: noqa: ANN003, FBT001, FBT002
|
# ruff: noqa: ANN003, FBT001, FBT002
|
||||||
|
|
||||||
|
|||||||
+398
-189
@@ -4,22 +4,24 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import importlib
|
import importlib
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
|
from functools import lru_cache
|
||||||
from sys import stderr
|
from sys import stderr
|
||||||
from typing import Any, Callable, NamedTuple
|
from typing import Any, Callable, NamedTuple
|
||||||
|
|
||||||
import torch
|
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 comfy.samplers import KSampler, k_diffusion_sampling
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from tqdm.auto import trange
|
from tqdm.auto import trange
|
||||||
|
|
||||||
from . import noise
|
from . import noise, utils
|
||||||
|
|
||||||
|
|
||||||
class HistoryType(Enum):
|
class HistoryType(Enum):
|
||||||
ZERO = auto()
|
ZERO = auto()
|
||||||
RAND = auto()
|
RAND = auto()
|
||||||
SAMPLE = auto()
|
SAMPLE = auto()
|
||||||
|
SAMPLE_NORM = auto()
|
||||||
|
|
||||||
|
|
||||||
class GuidanceType(Enum):
|
class GuidanceType(Enum):
|
||||||
@@ -35,15 +37,34 @@ class GuidanceConfig(NamedTuple):
|
|||||||
latent: Tensor | None = None
|
latent: Tensor | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class MomentumMode(Enum):
|
||||||
|
CLASSIC = auto()
|
||||||
|
NEW = auto()
|
||||||
|
DENOISED = auto()
|
||||||
|
|
||||||
|
|
||||||
class SonarConfig(NamedTuple):
|
class SonarConfig(NamedTuple):
|
||||||
momentum: float = 0.95
|
momentum: float = 0.95
|
||||||
momentum_hist: float = 0.75
|
momentum_hist: float = 0.75
|
||||||
direction: float = 1.0
|
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
|
init: HistoryType = HistoryType.ZERO
|
||||||
noise_type: noise.NoiseType | None = None
|
noise_type: noise.NoiseType | None = None
|
||||||
custom_noise: noise.CustomNoise | None = None
|
custom_noise: noise.CustomNoise | None = None
|
||||||
rand_init_noise_type: noise.NoiseType | None = None
|
rand_init_noise_type: noise.NoiseType | None = None
|
||||||
|
rand_init_noise_multiplier: float | int = 1.0
|
||||||
guidance: GuidanceConfig | None = None
|
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:
|
class SonarBase:
|
||||||
@@ -53,14 +74,69 @@ class SonarBase:
|
|||||||
self.history_d = None
|
self.history_d = None
|
||||||
self.cfg = cfg
|
self.cfg = cfg
|
||||||
self.noise_sampler = None
|
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(
|
def set_noise_sampler(
|
||||||
self,
|
self,
|
||||||
x: Tensor,
|
x: Tensor,
|
||||||
sigmas,
|
sigmas: Tensor,
|
||||||
noise_sampler: Callable | None,
|
noise_sampler: Callable | None,
|
||||||
seed: int | None = None,
|
seed: int | None = None,
|
||||||
):
|
) -> Callable:
|
||||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||||
if noise_sampler is not None and self.cfg.noise_type not in {
|
if noise_sampler is not None and self.cfg.noise_type not in {
|
||||||
None,
|
None,
|
||||||
@@ -90,17 +166,32 @@ class SonarBase:
|
|||||||
self.noise_sampler = noise_sampler
|
self.noise_sampler = noise_sampler
|
||||||
return noise_sampler
|
return noise_sampler
|
||||||
|
|
||||||
def init_hist_d(self, x: Tensor) -> None:
|
def init_hist_d(
|
||||||
if self.history_d is not None:
|
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
|
return
|
||||||
|
cfg = self.cfg
|
||||||
|
init = cfg.init
|
||||||
# memorize delta momentum
|
# memorize delta momentum
|
||||||
if self.cfg.init == HistoryType.ZERO:
|
if init == HistoryType.ZERO:
|
||||||
self.history_d = 0
|
self.history_d = None
|
||||||
elif self.cfg.init == HistoryType.SAMPLE:
|
elif init == HistoryType.SAMPLE:
|
||||||
self.history_d = x
|
self.history_d = (
|
||||||
elif self.cfg.init == HistoryType.RAND:
|
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(
|
ns = noise.get_noise_sampler(
|
||||||
self.cfg.rand_init_noise_type,
|
cfg.rand_init_noise_type,
|
||||||
x,
|
x,
|
||||||
None,
|
None,
|
||||||
None,
|
None,
|
||||||
@@ -109,31 +200,124 @@ class SonarBase:
|
|||||||
normalized=True,
|
normalized=True,
|
||||||
)
|
)
|
||||||
self.history_d = ns(None, None)
|
self.history_d = ns(None, None)
|
||||||
|
if cfg.rand_init_noise_multiplier != 1:
|
||||||
|
self.history_d *= cfg.rand_init_noise_multiplier
|
||||||
else:
|
else:
|
||||||
raise ValueError("Sonar sampler: bad history type")
|
raise ValueError("Sonar sampler: bad history type")
|
||||||
|
|
||||||
def update_hist(self, momentum_d):
|
@property
|
||||||
q = 1.0 - self.cfg.momentum_hist
|
@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
|
hd = self.history_d
|
||||||
if isinstance(hd, int) and hd == 0:
|
momentum_denoised = self.momentum_mix(
|
||||||
self.history_d = momentum_d
|
hd,
|
||||||
else:
|
denoised,
|
||||||
self.history_d = (1.0 - q) * hd + q * momentum_d
|
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):
|
def get_momentum_d(
|
||||||
if self.cfg.momentum == 1.0:
|
self,
|
||||||
return x + d * dt
|
x: Tensor,
|
||||||
|
denoised: Tensor,
|
||||||
|
sigma: Tensor,
|
||||||
|
*,
|
||||||
|
step: int,
|
||||||
|
momentum: float | None = None,
|
||||||
|
d: Tensor | None = None,
|
||||||
|
update_history=True,
|
||||||
|
) -> Tensor:
|
||||||
hd = self.history_d
|
hd = self.history_d
|
||||||
# correct current `d` with momentum
|
cfg = self.cfg
|
||||||
p = (1.0 - self.cfg.momentum) * self.cfg.direction
|
momentum = cfg.momentum if momentum is None else momentum
|
||||||
momentum_d = (1.0 - p) * d + p * hd
|
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
|
def momentum_step(
|
||||||
x = x + momentum_d * dt # noqa: PLR6104
|
self,
|
||||||
|
step: int,
|
||||||
self.update_hist(momentum_d)
|
x: Tensor,
|
||||||
|
denoised: Tensor,
|
||||||
return x
|
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:
|
class SonarGuidanceMixin:
|
||||||
@@ -152,17 +336,26 @@ class SonarGuidanceMixin:
|
|||||||
def prepare_ref_latent(latent: Tensor | None) -> Tensor:
|
def prepare_ref_latent(latent: Tensor | None) -> Tensor:
|
||||||
if latent is None:
|
if latent is None:
|
||||||
return None
|
return None
|
||||||
avg_s = latent.mean(dim=[2, 3], keepdim=True)
|
avg_s = latent.mean(dim=(-2, -1), keepdim=True)
|
||||||
std_s = latent.std(dim=[2, 3], keepdim=True)
|
std_s = latent.std(dim=(-2, -1), keepdim=True)
|
||||||
return ((latent - avg_s) / std_s).to(latent.dtype)
|
return (latent - avg_s).div_(std_s).to(latent.dtype)
|
||||||
|
|
||||||
def guidance_step(self, step_index: int, x: Tensor, denoised: Tensor):
|
def guidance_step(self, step_index: int, x: Tensor, denoised: Tensor) -> Tensor:
|
||||||
if self.guidance is None or self.guidance.factor == 0.0 or not self.guidance.start_step <= step_index <= self.guidance.end_step:
|
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
|
return x
|
||||||
if self.ref_latent.device != x.device:
|
if self.ref_latent.device != x.device:
|
||||||
self.ref_latent = self.ref_latent.to(device=x.device)
|
self.ref_latent = self.ref_latent.to(device=x.device)
|
||||||
if self.guidance.guidance_type == GuidanceType.LINEAR:
|
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:
|
if self.guidance.guidance_type == GuidanceType.EULER:
|
||||||
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
|
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
|
||||||
return self.guidance_euler(
|
return self.guidance_euler(
|
||||||
@@ -184,20 +377,26 @@ class SonarGuidanceMixin:
|
|||||||
ref_latent: Tensor,
|
ref_latent: Tensor,
|
||||||
factor: float = 0.2,
|
factor: float = 0.2,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
avg_t = denoised.mean(dim=[1, 2, 3], keepdim=True)
|
avg_t = denoised.mean(dim=(-3, -2, -1), keepdim=True)
|
||||||
std_t = denoised.std(dim=[1, 2, 3], keepdim=True)
|
std_t = denoised.std(dim=(-3, -2, -1), keepdim=True)
|
||||||
ref_img_shift = ref_latent * std_t + avg_t
|
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
|
dt = (sigma_next - sigma) * factor
|
||||||
return x + d * dt
|
return (d * dt).add_(x)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def guidance_linear(x: Tensor, ref_latent: Tensor, factor: float = 0.2) -> Tensor:
|
def guidance_linear(
|
||||||
avg_t = x.mean(dim=[1, 2, 3], keepdim=True)
|
x: Tensor,
|
||||||
std_t = x.std(dim=[1, 2, 3], keepdim=True)
|
ref_latent: Tensor,
|
||||||
ref_img_shift = ref_latent * std_t + avg_t
|
factor: float = 0.2,
|
||||||
return (1.0 - factor) * x + factor * ref_img_shift
|
*,
|
||||||
|
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):
|
class SonarWithGuidance(SonarBase, SonarGuidanceMixin):
|
||||||
@@ -210,9 +409,9 @@ class SonarSampler(SonarWithGuidance):
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
model,
|
model,
|
||||||
sigmas,
|
sigmas: Tensor,
|
||||||
s_in,
|
s_in: Tensor,
|
||||||
extra_args,
|
extra_args: dict[str, Any],
|
||||||
*args: list[Any],
|
*args: list[Any],
|
||||||
**kwargs: dict[str, Any],
|
**kwargs: dict[str, Any],
|
||||||
):
|
):
|
||||||
@@ -222,62 +421,49 @@ class SonarSampler(SonarWithGuidance):
|
|||||||
self.s_in = s_in
|
self.s_in = s_in
|
||||||
self.extra_args = extra_args
|
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):
|
class SonarEuler(SonarSampler):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
s_churn: float = 0.0,
|
|
||||||
s_tmin: float = 0.0,
|
|
||||||
s_tmax: float = float("inf"),
|
|
||||||
s_noise: float = 1.0,
|
|
||||||
*args: list[Any],
|
*args: list[Any],
|
||||||
**kwargs: dict[str, Any],
|
**kwargs: dict[str, Any],
|
||||||
):
|
):
|
||||||
super().__init__(*args, **kwargs)
|
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(
|
def step(self, step_index: int, sample: torch.FloatTensor):
|
||||||
self,
|
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
|
||||||
step_index: int,
|
|
||||||
sample: torch.FloatTensor,
|
|
||||||
):
|
|
||||||
self.init_hist_d(sample)
|
|
||||||
|
|
||||||
sigma, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1]
|
denoised = self.call_model(sample, sigma)
|
||||||
|
result_sample = self.momentum_step(
|
||||||
gamma = (
|
step_index,
|
||||||
min(self.s_churn / (len(self.sigmas) - 1), 2**0.5 - 1)
|
sample,
|
||||||
if self.s_tmin <= sigma <= self.s_tmax
|
denoised,
|
||||||
else 0.0
|
sigma,
|
||||||
|
sigma_next,
|
||||||
)
|
)
|
||||||
|
|
||||||
sigma_hat = sigma * (gamma + 1)
|
if sigma_next > 0:
|
||||||
|
|
||||||
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:
|
|
||||||
result_sample = self.guidance_step(step_index, result_sample, denoised)
|
result_sample = self.guidance_step(step_index, result_sample, denoised)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
result_sample,
|
result_sample,
|
||||||
sigma,
|
sigma,
|
||||||
sigma_hat,
|
sigma,
|
||||||
denoised,
|
denoised,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -286,26 +472,18 @@ class SonarEuler(SonarSampler):
|
|||||||
def sampler(
|
def sampler(
|
||||||
cls,
|
cls,
|
||||||
model,
|
model,
|
||||||
x,
|
x: Tensor,
|
||||||
sigmas,
|
sigmas: Tensor,
|
||||||
extra_args=None,
|
extra_args: dict | None = None,
|
||||||
callback=None,
|
callback=None,
|
||||||
disable=None,
|
disable: bool | None = None, # noqa: FBT001
|
||||||
noise_sampler: Callable | None = None,
|
noise_sampler: Callable | None = None,
|
||||||
sonar_config=None,
|
sonar_config: SonarConfig | None = None,
|
||||||
s_churn=0.0,
|
sonar_params: dict | None = None,
|
||||||
s_tmin=0.0,
|
) -> Tensor:
|
||||||
s_tmax=float("inf"),
|
sonar_config = cls.get_config(sonar_config, sonar_params)
|
||||||
s_noise=1.0,
|
s_in = x.new_ones((x.shape[0],))
|
||||||
):
|
|
||||||
if sonar_config is None:
|
|
||||||
sonar_config = SonarConfig()
|
|
||||||
s_in = x.new_ones([x.shape[0]])
|
|
||||||
sonar = cls(
|
sonar = cls(
|
||||||
s_churn,
|
|
||||||
s_tmin,
|
|
||||||
s_tmax,
|
|
||||||
s_noise,
|
|
||||||
model,
|
model,
|
||||||
sigmas,
|
sigmas,
|
||||||
s_in,
|
s_in,
|
||||||
@@ -320,7 +498,7 @@ class SonarEuler(SonarSampler):
|
|||||||
)
|
)
|
||||||
|
|
||||||
for i in trange(len(sigmas) - 1, disable=disable):
|
for i in trange(len(sigmas) - 1, disable=disable):
|
||||||
x, _sigma, sigma_hat, denoised = sonar.step(
|
x, sigma, sigma_hat, denoised = sonar.step(
|
||||||
i,
|
i,
|
||||||
x,
|
x,
|
||||||
)
|
)
|
||||||
@@ -329,7 +507,7 @@ class SonarEuler(SonarSampler):
|
|||||||
{
|
{
|
||||||
"x": x,
|
"x": x,
|
||||||
"i": i,
|
"i": i,
|
||||||
"sigma": sigmas[i],
|
"sigma": sigma,
|
||||||
"sigma_hat": sigma_hat,
|
"sigma_hat": sigma_hat,
|
||||||
"denoised": denoised,
|
"denoised": denoised,
|
||||||
},
|
},
|
||||||
@@ -354,31 +532,32 @@ class SonarEulerAncestral(SonarSampler):
|
|||||||
step_index: int,
|
step_index: int,
|
||||||
sample: torch.FloatTensor,
|
sample: torch.FloatTensor,
|
||||||
):
|
):
|
||||||
self.init_hist_d(sample)
|
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
|
||||||
|
sigma_down, sigma_up = get_ancestral_step(
|
||||||
sigma_from, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1]
|
sigma,
|
||||||
sigma_down, sigma_up = sampling.get_ancestral_step(
|
sigma_next,
|
||||||
sigma_from,
|
|
||||||
sigma_to,
|
|
||||||
eta=self.eta,
|
eta=self.eta,
|
||||||
)
|
)
|
||||||
|
|
||||||
denoised = self.model(sample, sigma_from * self.s_in, **self.extra_args)
|
denoised = self.call_model(sample, sigma)
|
||||||
derivative = sampling.to_d(sample, sigma_from, denoised)
|
result_sample = self.momentum_step(
|
||||||
dt = sigma_down - sigma_from
|
step_index,
|
||||||
|
sample,
|
||||||
result_sample = self.momentum_step(sample, derivative, dt)
|
denoised,
|
||||||
if sigma_to > 0:
|
sigma,
|
||||||
|
sigma_down,
|
||||||
|
)
|
||||||
|
if sigma_next > 0:
|
||||||
result_sample = self.guidance_step(step_index, result_sample, denoised)
|
result_sample = self.guidance_step(step_index, result_sample, denoised)
|
||||||
result_sample = ( # noqa: PLR6104
|
result_sample = ( # noqa: PLR6104
|
||||||
result_sample
|
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 (
|
return (
|
||||||
result_sample,
|
result_sample,
|
||||||
sigma_from,
|
sigma,
|
||||||
sigma_from,
|
sigma,
|
||||||
denoised,
|
denoised,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -392,14 +571,14 @@ class SonarEulerAncestral(SonarSampler):
|
|||||||
extra_args=None,
|
extra_args=None,
|
||||||
callback=None,
|
callback=None,
|
||||||
disable=None,
|
disable=None,
|
||||||
sonar_config=None,
|
sonar_config: SonarConfig | None = None,
|
||||||
|
sonar_params: dict | None = None,
|
||||||
eta=1.0,
|
eta=1.0,
|
||||||
s_noise=1.0,
|
s_noise=1.0,
|
||||||
noise_sampler: Callable | None = None,
|
noise_sampler: Callable | None = None,
|
||||||
):
|
):
|
||||||
if sonar_config is None:
|
sonar_config = cls.get_config(sonar_config, sonar_params)
|
||||||
sonar_config = SonarConfig()
|
s_in = x.new_ones((x.shape[0],))
|
||||||
s_in = x.new_ones([x.shape[0]])
|
|
||||||
sonar = cls(
|
sonar = cls(
|
||||||
eta,
|
eta,
|
||||||
s_noise,
|
s_noise,
|
||||||
@@ -449,104 +628,134 @@ class SonarDPMPPSDE(SonarSampler):
|
|||||||
self.s_noise = s_noise
|
self.s_noise = s_noise
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def sigma_fn(t) -> float:
|
def sigma_fn(t: Tensor) -> float:
|
||||||
return t.neg().exp()
|
return t.neg().exp()
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def t_fn(sigma) -> float:
|
def t_fn(sigma: Tensor) -> float:
|
||||||
return sigma.log.neg()
|
return sigma.log().neg()
|
||||||
|
|
||||||
# DPM++ solver algorithm copied from ComfyUI source.
|
# DPM++ solver algorithm copied from ComfyUI source.
|
||||||
def momentum_step( # noqa: PLR0914
|
def momentum_step( # noqa: PLR0914
|
||||||
self,
|
self,
|
||||||
step_index,
|
step_index: int,
|
||||||
x: Tensor,
|
x: Tensor,
|
||||||
denoised: Tensor,
|
denoised: Tensor,
|
||||||
sigma_from,
|
sigma: Tensor,
|
||||||
sigma_to,
|
sigma_next: Tensor,
|
||||||
sigma_down,
|
sigma_down: Tensor,
|
||||||
):
|
) -> Tensor:
|
||||||
if sigma_to == 0:
|
if sigma_next == 0:
|
||||||
derivative = sampling.to_d(x, sigma_from, denoised)
|
return super().momentum_step(step_index, x, denoised, sigma, sigma_down)
|
||||||
dt = sigma_down - sigma_from
|
|
||||||
return super().momentum_step(x, derivative, dt)
|
|
||||||
|
|
||||||
def sigma_fn(t):
|
cfg = self.cfg
|
||||||
return t.neg().exp()
|
# Halve the momentum proportion if there's history since we will use it twice.
|
||||||
|
adjusted_momentum = (
|
||||||
def t_fn(sigma):
|
cfg.momentum + (1 - cfg.momentum) / 2
|
||||||
return sigma.log().neg()
|
if self.history_d is not None
|
||||||
|
else cfg.momentum
|
||||||
hd = self.history_d
|
)
|
||||||
p = (1.0 - self.cfg.momentum) * self.cfg.direction
|
|
||||||
|
|
||||||
r = 1 / 2
|
r = 1 / 2
|
||||||
# DPM-Solver++
|
# 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
|
h = t_next - t
|
||||||
s = t + h * r
|
s = t + h * r
|
||||||
fac = 1 / (2 * r)
|
fac = 1 / (2 * r)
|
||||||
|
|
||||||
# Step 1
|
# Step 1
|
||||||
sd, su = sampling.get_ancestral_step(sigma_fn(t), sigma_fn(s), self.eta)
|
s_t, s_s = self.sigma_fn(t), self.sigma_fn(s)
|
||||||
s_ = t_fn(sd)
|
sd, su = get_ancestral_step(
|
||||||
diff_2 = (t - s_).expm1() * denoised
|
s_t,
|
||||||
momentum_d = (1.0 - p) * diff_2 + p * hd
|
s_s,
|
||||||
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),
|
|
||||||
self.eta,
|
self.eta,
|
||||||
)
|
)
|
||||||
t_next_ = t_fn(sd)
|
s_ = self.t_fn(sd)
|
||||||
denoised_d = (1 - fac) * denoised + fac * denoised_2
|
momentum_denoised = self.get_momentum_denoised(
|
||||||
diff_1 = (t - t_next_).expm1() * denoised_d
|
x,
|
||||||
momentum_d = (1.0 - p) * diff_1 + p * hd
|
denoised,
|
||||||
self.update_hist(momentum_d)
|
sigma,
|
||||||
x = (sigma_fn(t_next_) / sigma_fn(t)) * x - momentum_d
|
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)
|
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(
|
def step(
|
||||||
self,
|
self,
|
||||||
step_index: int,
|
step_index: int,
|
||||||
sample: torch.FloatTensor,
|
sample: torch.FloatTensor,
|
||||||
):
|
) -> Tensor:
|
||||||
def sigma_fn(t):
|
def sigma_fn(t):
|
||||||
return t.neg().exp()
|
return t.neg().exp()
|
||||||
|
|
||||||
def t_fn(sigma):
|
def t_fn(sigma):
|
||||||
return sigma.log().neg()
|
return sigma.log().neg()
|
||||||
|
|
||||||
self.init_hist_d(sample)
|
sigma, sigma_next = self.sigmas[step_index], self.sigmas[step_index + 1]
|
||||||
|
sigma_down, _sigma_up = get_ancestral_step(
|
||||||
sigma_from, sigma_to = self.sigmas[step_index], self.sigmas[step_index + 1]
|
sigma,
|
||||||
sigma_down, _sigma_up = sampling.get_ancestral_step(
|
sigma_next,
|
||||||
sigma_from,
|
|
||||||
sigma_to,
|
|
||||||
eta=self.eta,
|
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(
|
result_sample = self.momentum_step(
|
||||||
step_index,
|
step_index,
|
||||||
sample,
|
sample,
|
||||||
denoised,
|
denoised,
|
||||||
sigma_from,
|
sigma,
|
||||||
sigma_to,
|
sigma_next,
|
||||||
sigma_down,
|
sigma_down,
|
||||||
)
|
)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
result_sample,
|
result_sample,
|
||||||
sigma_from,
|
sigma,
|
||||||
sigma_from,
|
sigma,
|
||||||
denoised,
|
denoised,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -555,19 +764,19 @@ class SonarDPMPPSDE(SonarSampler):
|
|||||||
def sampler(
|
def sampler(
|
||||||
cls,
|
cls,
|
||||||
model,
|
model,
|
||||||
x,
|
x: Tensor,
|
||||||
sigmas,
|
sigmas: Tensor,
|
||||||
extra_args=None,
|
extra_args: dict | None = None,
|
||||||
callback=None,
|
callback=None,
|
||||||
disable=None,
|
disable: bool | None = None, # noqa: FBT001
|
||||||
sonar_config=None,
|
sonar_config: SonarConfig | None = None,
|
||||||
|
sonar_params: dict | None = None,
|
||||||
eta=1.0,
|
eta=1.0,
|
||||||
s_noise=1.0,
|
s_noise=1.0,
|
||||||
noise_sampler=None,
|
noise_sampler=None,
|
||||||
):
|
) -> Tensor:
|
||||||
if sonar_config is None:
|
sonar_config = cls.get_config(sonar_config, sonar_params)
|
||||||
sonar_config = SonarConfig()
|
s_in = x.new_ones((x.shape[0],))
|
||||||
s_in = x.new_ones([x.shape[0]])
|
|
||||||
sonar = cls(
|
sonar = cls(
|
||||||
eta,
|
eta,
|
||||||
s_noise,
|
s_noise,
|
||||||
@@ -602,7 +811,7 @@ class SonarDPMPPSDE(SonarSampler):
|
|||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
def add_samplers():
|
def add_samplers() -> None:
|
||||||
extra_samplers = {
|
extra_samplers = {
|
||||||
"sonar_euler": SonarEuler.sampler,
|
"sonar_euler": SonarEuler.sampler,
|
||||||
"sonar_euler_ancestral": SonarEulerAncestral.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