Sync changes
This commit is contained in:
+7
-1
@@ -64,6 +64,7 @@ class SonarLatentOperationAdvanced(SonarLatentOperation):
|
||||
*,
|
||||
blend_mode: str,
|
||||
blend_strength: float,
|
||||
blend_strategy: str,
|
||||
input_multiplier: float,
|
||||
output_multiplier: float,
|
||||
difference_multiplier: float,
|
||||
@@ -74,6 +75,7 @@ class SonarLatentOperationAdvanced(SonarLatentOperation):
|
||||
super().__init__(**kwargs)
|
||||
self.blend_function = utils.BLENDING_MODES[blend_mode]
|
||||
self.blend_strength = blend_strength
|
||||
self.blend_strategy = blend_strategy
|
||||
self.input_multiplier = input_multiplier
|
||||
self.output_multiplier = output_multiplier
|
||||
self.difference_multiplier = difference_multiplier
|
||||
@@ -103,7 +105,11 @@ class SonarLatentOperationAdvanced(SonarLatentOperation):
|
||||
) - t
|
||||
if self.difference_multiplier != 1.0:
|
||||
diff *= self.difference_multiplier
|
||||
return self.blend_function(t, diff, self.blend_strength)
|
||||
if self.blend_strategy == "difference":
|
||||
return self.blend_function(t, diff, self.blend_strength)
|
||||
if self.blend_strategy == "result":
|
||||
return self.blend_function(t, t + diff, self.blend_strength)
|
||||
raise ValueError(f"Unknown blend strategy: {self.blend_strategy}")
|
||||
|
||||
|
||||
class SonarLatentOperationNoise(SonarLatentOperation):
|
||||
|
||||
@@ -361,10 +361,6 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
|
||||
|
||||
INPUT_TYPES = SonarLazyInputTypes(
|
||||
lambda: SonarInputTypes()
|
||||
.req_field_operation(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Latent operation to apply.",
|
||||
)
|
||||
.req_float_start_sigma(
|
||||
default=-1.0,
|
||||
min=-1.0,
|
||||
@@ -395,6 +391,15 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
|
||||
default=0.5,
|
||||
tooltip="Strength of the blend.",
|
||||
)
|
||||
.req_field_blend_strategy(
|
||||
("difference", "result"),
|
||||
default="difference",
|
||||
tooltip="Controls whether blending occurs with the difference or changed result after the latent operation.",
|
||||
)
|
||||
.opt_field_operation(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Latent operation to apply.",
|
||||
)
|
||||
.opt_field_operation_alt(
|
||||
"LATENT_OPERATION",
|
||||
tooltip="Optional alternative operation that will be used when the primary one isn't enabled. May be useful in a case when you want one operation between sigma 1.0 and 0.5 and then a difference operation for lower sigmas which is kind of annoying to specify manually (you'd need to do something like configure another operation to start at 0.499999 or something).",
|
||||
@@ -421,7 +426,6 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
|
||||
def go(
|
||||
cls,
|
||||
*,
|
||||
operation,
|
||||
start_sigma: float,
|
||||
end_sigma: float,
|
||||
input_multiplier: float,
|
||||
@@ -429,6 +433,8 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
|
||||
difference_multiplier: float,
|
||||
blend_mode: str,
|
||||
blend_strength: float,
|
||||
blend_strategy: str,
|
||||
operation=None,
|
||||
operation_alt=None,
|
||||
operation_2=None,
|
||||
operation_3=None,
|
||||
@@ -456,6 +462,7 @@ class SonarLatentOperationAdvancedNode(metaclass=IntegratedNode):
|
||||
difference_multiplier=difference_multiplier,
|
||||
blend_mode=blend_mode,
|
||||
blend_strength=blend_strength,
|
||||
blend_strategy=blend_strategy,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
+57
-14
@@ -4,12 +4,13 @@ import functools
|
||||
import inspect
|
||||
import math
|
||||
import random
|
||||
from typing import Any, Callable
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import yaml
|
||||
from comfy import model_management, samplers
|
||||
from comfy import utils as comfy_utils
|
||||
from tqdm import tqdm
|
||||
|
||||
from .. import noise, utils
|
||||
@@ -24,6 +25,14 @@ from .base import (
|
||||
SonarNormalizeNoiseNodeMixin,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
try:
|
||||
from comfy import nested_tensor
|
||||
except (ModuleNotFoundError, ImportError):
|
||||
nested_tensor = None
|
||||
|
||||
|
||||
class NoisyLatentLikeNode(metaclass=IntegratedNode):
|
||||
DESCRIPTION = "Allows generating noise (and optionally adding it) based on a reference latent. Note: For img2img workflows, you will generally want to enable add_to_latent as well as connecting the model and sigmas inputs."
|
||||
@@ -390,10 +399,29 @@ class CustomNOISE:
|
||||
return result.to_sparse()
|
||||
errstr = f"Cannot handle latent layout {type(latent_image.layout).__name__}"
|
||||
raise NotImplementedError(errstr)
|
||||
return result if self.multiplier == 1.0 else result.mul_(self.multiplier)
|
||||
if self.multiplier != 1.0:
|
||||
result *= self.multiplier
|
||||
return result
|
||||
|
||||
def generate_noise(self, input_latent):
|
||||
latent_image = input_latent["samples"]
|
||||
latent_image = orig_latent_image = input_latent["samples"]
|
||||
orig_type = type(latent_image)
|
||||
if nested_tensor is not None and latent_image.is_nested:
|
||||
latent_image, latent_shapes = comfy_utils.pack_latents(
|
||||
latent_image.unbind(),
|
||||
)
|
||||
else:
|
||||
latent_shapes = None
|
||||
|
||||
def result_out(result, latent_shapes):
|
||||
if latent_shapes is None:
|
||||
return result
|
||||
tensors = comfy_utils.unpack_latents(result, latent_shapes)
|
||||
if isinstance(orig_type, torch.Tensor):
|
||||
# Should be an actual PyTorch nested tensor, not ComfyUI's custom class
|
||||
return orig_type(tensors, layout=orig_latent_image.layout)
|
||||
return orig_type(tensors)
|
||||
# TODO: Test this stuff.
|
||||
batch_inds = input_latent.get("batch_index")
|
||||
torch.manual_seed(self.seed)
|
||||
random.seed(self.seed)
|
||||
@@ -405,18 +433,33 @@ class CustomNOISE:
|
||||
device="cpu",
|
||||
)
|
||||
if batch_inds is None:
|
||||
return self._sample_noise(latent_image, self.seed)
|
||||
unique_inds, inverse_inds = np.unique(batch_inds, return_inverse=True)
|
||||
result = []
|
||||
batch_size = latent_image.shape[0]
|
||||
for idx in range(unique_inds[-1] + 1):
|
||||
noise = self._sample_noise(
|
||||
latent_image[idx % batch_size].unsqueeze(0),
|
||||
self.seed + idx,
|
||||
return result_out(
|
||||
self._sample_noise(latent_image, self.seed),
|
||||
latent_shapes,
|
||||
)
|
||||
if idx in unique_inds:
|
||||
result.append(noise)
|
||||
return torch.cat(tuple(result[i] for i in inverse_inds), axis=0)
|
||||
batch_size = latent_image.shape[0]
|
||||
unique_inds, inverse_inds = np.unique(batch_inds, return_inverse=True)
|
||||
use_idxs = {
|
||||
out_idx: idx % batch_size for out_idx, idx in enumerate(unique_inds)
|
||||
}
|
||||
use_idxs = (idx for idx in range(unique_inds[-1] + 1) if idx in unique_inds)
|
||||
use_idxs = {idx: inverse_inds[uidx] for uidx, idx in enumerate(use_idxs)}
|
||||
result = torch.empty(
|
||||
(len(use_idxs), *latent_image.shape[1:]),
|
||||
dtype=latent_image.dtype,
|
||||
device=latent_image.device,
|
||||
)
|
||||
for idx in range(unique_inds[-1] + 1):
|
||||
sample_idx = idx % batch_size
|
||||
sample = latent_image[sample_idx].unsqueeze(0)
|
||||
print(
|
||||
f"\nNOISE: idx {idx}, sample_idx {sample_idx}, shape {latent_image[sample_idx].shape}, nested={sample.is_nested}",
|
||||
)
|
||||
noise = self._sample_noise(sample, self.seed + idx)
|
||||
batch_out_idx = use_idxs.get(idx)
|
||||
if batch_out_idx is not None:
|
||||
result[batch_out_idx : batch_out_idx + 1] = noise[:1]
|
||||
return result_out(result, latent_shapes)
|
||||
|
||||
|
||||
class SonarToComfyNOISENode(metaclass=IntegratedNode):
|
||||
|
||||
+800
-675
File diff suppressed because it is too large
Load Diff
@@ -755,12 +755,13 @@ class SonarAdvancedSimulationNoiseNode(SonarCustomNoiseNodeBase):
|
||||
.req_field_band_shape(
|
||||
("log_gaussian", "raised_cosine"),
|
||||
default="log_gaussian",
|
||||
tooltip="TBD",
|
||||
tooltip="No effect in power_law spectral mode or in multi_octave spectral mode when octaves is set to 0.",
|
||||
)
|
||||
.req_field_channel_mode(
|
||||
(
|
||||
"stacked",
|
||||
"over_depth",
|
||||
"flat",
|
||||
"over_depth_alt",
|
||||
"over_depth_avg",
|
||||
"over_depth_h",
|
||||
@@ -811,8 +812,8 @@ class SonarAdvancedSimulationNoiseNode(SonarCustomNoiseNodeBase):
|
||||
)
|
||||
.req_int_octaves(
|
||||
default=3,
|
||||
min=1,
|
||||
tooltip="Number of octaves of noise to generate. Only has an effect in multi_octave spectral mode.",
|
||||
min=0,
|
||||
tooltip="Number of octaves of noise to generate. Only has an effect in multi_octave spectral mode. You can also set octaves to 0 to disable octaves.",
|
||||
)
|
||||
.req_float_lacunarity(default=2.0)
|
||||
.req_float_gain(default=0.75)
|
||||
|
||||
+90
-18
@@ -10,6 +10,7 @@ import comfy
|
||||
import torch
|
||||
import yaml
|
||||
from comfy.k_diffusion import sampling
|
||||
from comfy.model_management import throw_exception_if_processing_interrupted
|
||||
from torch import Tensor
|
||||
|
||||
from . import external, utils
|
||||
@@ -64,7 +65,16 @@ class CustomNoiseItemBase(abc.ABC):
|
||||
|
||||
def get_normalize(self, k, default=None):
|
||||
val = getattr(self, k, None)
|
||||
return default if val is None else val
|
||||
if val in {None, "default"}:
|
||||
return default
|
||||
if val == "disabled":
|
||||
return False
|
||||
if val == "forced":
|
||||
return True
|
||||
return default
|
||||
# if isinstance(val, bool):
|
||||
# return val
|
||||
# return default if val is None else val
|
||||
|
||||
@abc.abstractmethod
|
||||
def make_noise_sampler(
|
||||
@@ -1691,6 +1701,62 @@ class ScatternetFilteredNoise(CustomNoiseItemBase):
|
||||
return noise_sampler
|
||||
|
||||
|
||||
class NoveltyFilteredNoise(CustomNoiseItemBase):
|
||||
def clone_key(self, k):
|
||||
if k == "noise" and self.noise is not None:
|
||||
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)
|
||||
if self.noise is not None:
|
||||
internal_ns = self.noise.make_noise_sampler(
|
||||
x,
|
||||
*args,
|
||||
sigma_min=sigma_min,
|
||||
sigma_max=sigma_max,
|
||||
normalized=self.normalize_noise,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
internal_ns = None
|
||||
ns_kwargs = getattr(self, "ns_kwargs", {}).copy()
|
||||
kwargs |= ns_kwargs
|
||||
ns = NoveltyFilteredNoiseGenerator(
|
||||
x,
|
||||
*args,
|
||||
sigma_min=sigma_min,
|
||||
sigma_max=sigma_max,
|
||||
normalized=False,
|
||||
noise_sampler=internal_ns,
|
||||
skip_initial=self.skip_initial,
|
||||
iters_per_call=self.iters_per_call,
|
||||
blend_ratio=self.blend_ratio,
|
||||
blend_function=utils.BLENDING_MODES[self.blend_mode],
|
||||
update_blend_ratio=self.update_blend_ratio,
|
||||
update_blend_function=utils.BLENDING_MODES[self.update_blend_mode],
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def noise_sampler(sigma, sigma_next):
|
||||
return scale_noise(
|
||||
ns(sigma, sigma_next),
|
||||
factor,
|
||||
normalized=normalize,
|
||||
)
|
||||
|
||||
return noise_sampler
|
||||
|
||||
|
||||
class LatentOperationFilteredNoise(CustomNoiseItemBase):
|
||||
def clone_key(self, k):
|
||||
if k == "noise" and self.noise is not None:
|
||||
@@ -1895,29 +1961,32 @@ class PerDimNoise(CustomNoiseItemBase):
|
||||
slice(-dim_size, None) if d == dim else slice(None, None)
|
||||
for d in range(x.ndim)
|
||||
)
|
||||
n_chunks = math.ceil(dim_size / chunk_size)
|
||||
|
||||
if self.shrink_dim:
|
||||
|
||||
def noise_sampler(sigma, sigma_next) -> torch.Tensor:
|
||||
noise = torch.cat(
|
||||
tuple(ns(sigma, sigma_next) for _ in range(dim_size)),
|
||||
dim=dim,
|
||||
)[trim_slice]
|
||||
chunks = []
|
||||
for _ in range(n_chunks):
|
||||
throw_exception_if_processing_interrupted()
|
||||
chunks.append(ns(sigma, sigma_next))
|
||||
noise = torch.cat(chunks, dim=dim)[trim_slice]
|
||||
return scale_noise(noise, factor, normalized=normalize)
|
||||
|
||||
else:
|
||||
select_dim = [slice(None, None) for d in range(x.ndim)]
|
||||
n_chunks = math.ceil(dim_size / chunk_size)
|
||||
temp_shape = list(x.shape)
|
||||
temp_shape[dim] = int(n_chunks * chunk_size)
|
||||
return noise_sampler
|
||||
|
||||
def noise_sampler(sigma, sigma_next) -> torch.Tensor:
|
||||
nonlocal select_dim
|
||||
result = x.new_zeros(temp_shape)
|
||||
# result = torch.zeros_like(x)
|
||||
for idx in range(0, dim_size, chunk_size):
|
||||
select_dim[dim] = slice(idx, idx + chunk_size)
|
||||
result[select_dim] = ns(sigma, sigma_next)[select_dim]
|
||||
return scale_noise(result[trim_slice], factor, normalized=normalize)
|
||||
select_dim = [slice(None, None) for d in range(x.ndim)]
|
||||
temp_shape = list(x.shape)
|
||||
temp_shape[dim] = int(n_chunks * chunk_size)
|
||||
|
||||
def noise_sampler(sigma, sigma_next) -> torch.Tensor:
|
||||
nonlocal select_dim
|
||||
result = x.new_zeros(temp_shape)
|
||||
for idx in range(0, dim_size, chunk_size):
|
||||
throw_exception_if_processing_interrupted()
|
||||
select_dim[dim] = slice(idx, idx + chunk_size)
|
||||
result[select_dim] = ns(sigma, sigma_next)[select_dim]
|
||||
return scale_noise(result[trim_slice], factor, normalized=normalize)
|
||||
|
||||
return noise_sampler
|
||||
|
||||
@@ -2123,6 +2192,9 @@ class CustomNoiseParametersNoise(CustomNoiseItemBase):
|
||||
):
|
||||
factor = self.factor
|
||||
normalize = self.get_normalize("normalize", normalized)
|
||||
print(
|
||||
f"\n****** NS: normalized={normalized}, normalize={normalize}, self.normalize={self.normalize}"
|
||||
)
|
||||
orig_shape = x.shape
|
||||
orig_dtype = x.dtype
|
||||
orig_device = x.device
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from .base import MixedNoiseGenerator, NoiseError, NoiseType
|
||||
from .collatz_noise_generator import CollatzNoiseGenerator
|
||||
from .distro_noise_generator import DistroNoiseGenerator
|
||||
from .novelty_filtered_noise import NoveltyFilteredNoiseGenerator
|
||||
from .scatternet_filtered_noise_generator import ScatternetFilteredNoiseGenerator
|
||||
from .simple_noise_generators import (
|
||||
BrownianNoiseGenerator,
|
||||
@@ -34,6 +35,7 @@ __all__ = (
|
||||
"MixedNoiseGenerator",
|
||||
"NoiseError",
|
||||
"NoiseType",
|
||||
"NoveltyFilteredNoiseGenerator",
|
||||
"OneFNoiseGenerator",
|
||||
"PerlinOldNoiseGenerator",
|
||||
"PinkOldNoiseGenerator",
|
||||
|
||||
@@ -20,8 +20,6 @@ F = torch.nn.functional
|
||||
class CollatzNoiseGenerator(NoiseGenerator):
|
||||
name = "collatz"
|
||||
|
||||
chain_cache: ClassVar[dict] = {}
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
@@ -54,10 +52,10 @@ class CollatzNoiseGenerator(NoiseGenerator):
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _get_iter_slices(n_dims, dim, idx, stride) -> list:
|
||||
def _get_iter_slices(n_dims, dim, idx, stride) -> tuple:
|
||||
result = [slice(None)] * n_dims
|
||||
result[dim] = slice(idx, None, stride)
|
||||
return result
|
||||
return tuple(result)
|
||||
|
||||
def _generate_iteration(
|
||||
self,
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
# ruff: noqa: ANN002, ANN003
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from .base import NoiseGenerator
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
|
||||
def sum_rms_blend(
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
t: torch.Tensor | float = 1.0,
|
||||
*,
|
||||
orig_shape: torch.Size | tuple[int, ...],
|
||||
dims_a: tuple[int, ...] = (1,),
|
||||
dims_b: tuple[int, ...] = (-1, -2),
|
||||
) -> torch.Tensor:
|
||||
rms_a = a / math.prod(orig_shape[d] for d in dims_a) ** 0.5
|
||||
rms_b = b / math.prod(orig_shape[d] for d in dims_b) ** 0.5
|
||||
variance_a = rms_a.pow_(2.0)
|
||||
variance_b = rms_b.pow_(2.0)
|
||||
result = variance_a
|
||||
result += variance_b * t
|
||||
result /= 1.0 + t
|
||||
result **= 0.5
|
||||
return result
|
||||
|
||||
|
||||
def metrics_blend(
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
t: torch.Tensor | float = 1.0,
|
||||
*,
|
||||
orig_shape: torch.Size | tuple[int, ...],
|
||||
dims_a: tuple[int, ...] = (-1, -2),
|
||||
dims_b: tuple[int, ...] = (1,),
|
||||
use_rms: bool = True,
|
||||
rms_power: float = 2.0,
|
||||
) -> torch.Tensor:
|
||||
count_a = math.prod(orig_shape[d] for d in dims_a)
|
||||
count_b = math.prod(orig_shape[d] for d in dims_b)
|
||||
denom_a = count_a ** (1 / rms_power) if use_rms else count_a
|
||||
denom_b = count_b ** (1 / rms_power) if use_rms else count_b
|
||||
curr_a = a / denom_a
|
||||
curr_b = b / denom_b
|
||||
if use_rms:
|
||||
curr_a **= rms_power
|
||||
curr_b = curr_b.pow_(rms_power) * t
|
||||
result = curr_b.add_(curr_a)
|
||||
result /= 1.0 + t
|
||||
return result.pow_(1.0 / rms_power) if use_rms else result
|
||||
|
||||
|
||||
class NoveltyFilteredNoiseGenerator(NoiseGenerator):
|
||||
name = "novelty"
|
||||
|
||||
initial_noise_state: torch.Tensor | None = None
|
||||
noise_state: torch.Tensor | None = None
|
||||
blend_function: Callable | None = None
|
||||
|
||||
@classmethod
|
||||
def ng_params(cls):
|
||||
return super().ng_params() | {
|
||||
"skip_initial": 1,
|
||||
"iters_per_call": 1,
|
||||
"dim_groups": ((1,), (-1, -2)),
|
||||
"blend_ratio": 1.0,
|
||||
"blend_function": None,
|
||||
"update_blend_ratio": 1.0,
|
||||
"update_blend_function": None,
|
||||
"noise_sampler": None,
|
||||
}
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
if self.blend_function is None:
|
||||
raise ValueError("Missing blend function!")
|
||||
|
||||
def generate(self, *args) -> torch.Tensor:
|
||||
ng = (
|
||||
partial(self.noise_sampler, *args) if self.noise_sampler else self.rand_like
|
||||
)
|
||||
noise_state = self.noise_state
|
||||
had_state = self.noise_state is not None
|
||||
it_counter = 0 if had_state else 0 - self.skip_initial
|
||||
its_call = max(1, self.iters_per_call)
|
||||
bf = self.blend_function
|
||||
blend_ratio = self.blend_ratio
|
||||
update_blend_ratio = self.update_blend_ratio
|
||||
ubf = self.update_blend_function
|
||||
if ubf is None:
|
||||
# Linear weighted average
|
||||
def ubf(a: torch.Tensor, b: torch.Tensor, t: float) -> torch.Tensor:
|
||||
return (b * t).add_(a).div_(1.0 + abs(t))
|
||||
|
||||
call_initial_noise = None
|
||||
curr_noise = None
|
||||
call_noise_state = None
|
||||
while it_counter < its_call:
|
||||
if noise_state is None:
|
||||
noise_state = ng()
|
||||
self.initial_noise_state = noise_state.clone()
|
||||
self.noise_state = noise_state.clone()
|
||||
continue
|
||||
curr_noise = ng()
|
||||
if call_initial_noise is None:
|
||||
call_initial_noise = curr_noise.clone()
|
||||
seen = {id(curr_noise)}
|
||||
for ortho_target in (
|
||||
self.initial_noise_state,
|
||||
call_initial_noise if it_counter > 0 else None,
|
||||
noise_state,
|
||||
):
|
||||
tid = id(ortho_target)
|
||||
if ortho_target is None or tid in seen:
|
||||
continue
|
||||
seen.add(tid)
|
||||
curr_noise = bf(ortho_target, curr_noise, blend_ratio)
|
||||
curr_noise -= ortho_target
|
||||
it_counter += 1
|
||||
if it_counter < 1:
|
||||
self.noise_state = curr_noise.clone()
|
||||
noise_state = curr_noise
|
||||
continue
|
||||
if call_noise_state is None:
|
||||
call_noise_state = curr_noise
|
||||
else:
|
||||
call_noise_state = ubf(call_noise_state, curr_noise, update_blend_ratio)
|
||||
if call_noise_state is None:
|
||||
raise RuntimeError("Unexpected unpopulated call_noise_state!")
|
||||
# self.noise_state = call_noise_state.clone()
|
||||
self.noise_state = ubf(noise_state, call_noise_state, update_blend_ratio)
|
||||
return call_noise_state
|
||||
|
||||
# def generate(self, *args) -> torch.Tensor:
|
||||
# ng = (
|
||||
# partial(self.noise_sampler, *args) if self.noise_sampler else self.rand_like
|
||||
# )
|
||||
# noise_state = self.noise_state
|
||||
# had_state = self.noise_state is not None
|
||||
# it_counter = 0 if had_state else 0 - self.skip_initial
|
||||
# its_call = self.iters_per_call
|
||||
# bf = self.blend_function
|
||||
# blend_ratio = self.blend_ratio
|
||||
# update_blend_ratio = self.update_blend_ratio
|
||||
# ubf = self.update_blend_function
|
||||
# if ubf is None or True:
|
||||
# # Linear weighted average
|
||||
# def ubf(a: torch.Tensor, b: torch.Tensor, t: float) -> torch.Tensor:
|
||||
# return (b * t).add_(a).div_(1.0 + abs(t))
|
||||
|
||||
# call_initial_noise = None
|
||||
# while it_counter < its_call:
|
||||
# if noise_state is None:
|
||||
# noise_state = ng()
|
||||
# self.initial_noise_state = noise_state.clone()
|
||||
# continue
|
||||
# curr_noise = ng()
|
||||
# it_counter += 1
|
||||
# if it_counter < 1:
|
||||
# noise_state = curr_noise
|
||||
# continue
|
||||
# if call_initial_noise is None and self.iters_per_call > 1:
|
||||
# call_initial_noise = curr_noise.clone()
|
||||
# for ortho_target in (
|
||||
# self.initial_noise_state,
|
||||
# call_initial_noise if it_counter > 0 else None,
|
||||
# noise_state,
|
||||
# ):
|
||||
# if ortho_target is None:
|
||||
# continue
|
||||
# curr_noise = bf(ortho_target, curr_noise, blend_ratio)
|
||||
# curr_noise -= ortho_target
|
||||
# # curr_noise = bf(noise_state, ng(), blend_ratio).sub_(noise_state)
|
||||
# noise_state = ubf(noise_state, curr_noise, update_blend_ratio)
|
||||
# self.noise_state = noise_state.clone()
|
||||
# return noise_state
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
|
||||
import itertools
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
@@ -76,16 +77,20 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
torch.complex128 if self.dtype == torch.float64 else torch.complex64
|
||||
)
|
||||
self.eff_batch = (
|
||||
self.batch if cm != "stacked" else self.batch * math.ceil(self.channels / 3)
|
||||
self.batch
|
||||
if cm not in {"stacked", "flat"}
|
||||
else self.batch * math.ceil(self.channels / 3)
|
||||
)
|
||||
ns_shape = torch.Size(
|
||||
(
|
||||
self.eff_batch,
|
||||
self.depth * self.depth_increment,
|
||||
self.height,
|
||||
self.width,
|
||||
)
|
||||
)
|
||||
ns_shape = torch.Size((
|
||||
self.eff_batch,
|
||||
self.depth * self.depth_increment,
|
||||
self.height,
|
||||
self.width,
|
||||
))
|
||||
|
||||
def gaussian_noise_sampler(*_args: list) -> torch.Tensor:
|
||||
def gaussian_noise_sampler(*_args: Any) -> torch.Tensor:
|
||||
return torch.randn(ns_shape, dtype=self.cdtype, device=self.gen_device).to(
|
||||
device=self.device,
|
||||
)
|
||||
@@ -121,7 +126,7 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _radial_k(*ks: list) -> torch.Tensor:
|
||||
def _radial_k(*ks: torch.Tensor) -> torch.Tensor:
|
||||
"""Calculates the radial distance in k-space."""
|
||||
return sum(kt**2 for kt in ks).sqrt_()
|
||||
|
||||
@@ -187,6 +192,9 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
return wk_out(
|
||||
self._handle_band_shape(k_rad, self.band_pass_low, self.band_pass_high),
|
||||
)
|
||||
if self.octaves == 0:
|
||||
# Ones where k_rad is non-zero, otherwise zero.
|
||||
return (k_rad != 0).to(k_rad)
|
||||
base_k = 2 * math.pi / max(1, min(sizes)) if self.base_k == 0 else self.base_k
|
||||
wk = torch.zeros_like(k_rad)
|
||||
for o in range(self.octaves):
|
||||
@@ -205,9 +213,17 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
ns_args: tuple | list,
|
||||
**_kwargs,
|
||||
):
|
||||
n_dims = len(k_grids_orig)
|
||||
n_samplers = len(self.noise_samplers)
|
||||
f_fs = tuple(
|
||||
ns(*ns_args).to(device=self.device).mul_(wk) for ns in self.noise_samplers
|
||||
self.noise_samplers[ns_idx % n_samplers](*ns_args)
|
||||
.to(
|
||||
device=self.device,
|
||||
)
|
||||
.mul_(wk)
|
||||
for ns_idx in range(n_dims)
|
||||
)
|
||||
|
||||
# --- Perform the Helmholtz projection using the UN SCALED grids ---
|
||||
k_sq_proj = self._radial_k(*k_grids_orig) ** 2
|
||||
k_dot_f = sum(k_p * f_f for k_p, f_f in zip(k_grids_orig, f_fs))
|
||||
@@ -218,35 +234,6 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
f_f - k_grid * k_grid_scale for f_f, k_grid in zip(f_fs, k_grids_orig)
|
||||
)
|
||||
|
||||
# def _handle_field_curl(
|
||||
# self,
|
||||
# *,
|
||||
# k_rad: torch.Tensor,
|
||||
# k_grids_orig: tuple,
|
||||
# wk: torch.Tensor,
|
||||
# ns_args: tuple | list,
|
||||
# **_kwargs,
|
||||
# ):
|
||||
# if len(k_grids_orig) != 3:
|
||||
# raise ValueError("Can only handle 3 dimensions currently")
|
||||
# # inv_k_rad = torch.where(k_rad == 0, 0.0, 1.0 / k_rad)
|
||||
# # wk_potential = wk * inv_k_rad
|
||||
# wk_potential = wk / k_rad
|
||||
# f_fs = tuple(
|
||||
# ns(*ns_args).to(device=self.device).mul_(wk_potential)
|
||||
# for ns in self.noise_samplers
|
||||
# )
|
||||
# i_k_grids = tuple(
|
||||
# (1j * k_grid).to(dtype=self.cdtype) for k_grid in k_grids_orig
|
||||
# )
|
||||
# fz, fy, fx = f_fs
|
||||
# ikz, iky, ikx = i_k_grids
|
||||
# return (
|
||||
# ikx * fy - iky * fx, # z
|
||||
# ikz * fx - ikx * fz, # y
|
||||
# iky * fz - ikz * fy, # x
|
||||
# )
|
||||
|
||||
def _handle_field_curl(
|
||||
self,
|
||||
*,
|
||||
@@ -324,70 +311,6 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
|
||||
_handle_field_curl_ndim = _handle_field_curl
|
||||
|
||||
def _handle_field_basis_orig(
|
||||
self,
|
||||
*,
|
||||
k_rad: torch.Tensor,
|
||||
k_grids_orig: tuple,
|
||||
wk: torch.Tensor,
|
||||
ns_args: tuple | list,
|
||||
**_kwargs,
|
||||
) -> tuple:
|
||||
if len(k_grids_orig) != 3:
|
||||
raise ValueError("Can only handle 3 dimensions currently")
|
||||
k_norm_z, k_norm_y, k_norm_x = (
|
||||
torch.where(k_rad == 0, 0.0, k / k_rad) for k in k_grids_orig
|
||||
)
|
||||
|
||||
# Pick a fixed vector `ez` to take a cross product with.
|
||||
# Handle the singularity where k is parallel to ez.
|
||||
ez = torch.tensor((0.0, 0.0, 1.0), device=self.device, dtype=self.dtype)
|
||||
is_parallel = (k_norm_x.abs() < 1e-6) & (k_norm_y.abs() < 1e-6)
|
||||
|
||||
# First basis vector u = k x ez (or k x ey for the singularity)
|
||||
ux = torch.where(
|
||||
is_parallel,
|
||||
k_norm_y * 0 - k_norm_z * 1,
|
||||
k_norm_y * ez[2] - k_norm_z * ez[1],
|
||||
)
|
||||
uy = torch.where(
|
||||
is_parallel,
|
||||
k_norm_z * 0 - k_norm_x * 0,
|
||||
k_norm_z * ez[0] - k_norm_x * ez[2],
|
||||
)
|
||||
uz = torch.where(
|
||||
is_parallel,
|
||||
k_norm_x * 1 - k_norm_y * 0,
|
||||
k_norm_x * ez[1] - k_norm_y * ez[0],
|
||||
)
|
||||
|
||||
u_mag = torch.sqrt(sum((ux**2, uy**2, uz**2)))
|
||||
inv_u_mag = torch.where(u_mag == 0, 0.0, 1.0 / u_mag)
|
||||
ux *= inv_u_mag
|
||||
uy *= inv_u_mag
|
||||
uz *= inv_u_mag
|
||||
ux[k_rad == 0], uy[k_rad == 0], uz[k_rad == 0] = 0, 0, 0
|
||||
|
||||
# Second basis vector v = k x u
|
||||
vx = k_norm_y * uz - k_norm_z * uy
|
||||
vy = k_norm_z * ux - k_norm_x * uz
|
||||
vz = k_norm_x * uy - k_norm_y * ux
|
||||
|
||||
# 3. Generate two independent random complex scalar fields
|
||||
a_f = self.noise_samplers[0](*ns_args).to(device=self.device)
|
||||
b_f = self.noise_samplers[1](*ns_args).to(device=self.device)
|
||||
|
||||
# 4. Modulate the random fields by the spectral envelope
|
||||
a_f *= wk
|
||||
b_f *= wk
|
||||
|
||||
# 5. Project the random fields onto the basis vectors to form the final field
|
||||
return (
|
||||
a_f * uz + b_f * vz, # z
|
||||
a_f * uy + b_f * vy, # y
|
||||
a_f * ux + b_f * vx, # x
|
||||
)
|
||||
|
||||
def _handle_field_basis(
|
||||
self,
|
||||
*,
|
||||
@@ -407,9 +330,6 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
|
||||
# --- Case 1: 2D (simple and fast) ---
|
||||
if n_dims * nd_fixup == 2:
|
||||
if len(self.noise_samplers) < 1:
|
||||
raise ValueError("2D basis mode requires at least 1 noise sampler.")
|
||||
|
||||
# The basis is a single vector perpendicular to k: u = (-ky, kx)
|
||||
kn_y, kn_x = k_norm_components
|
||||
basis_vectors = [
|
||||
@@ -420,6 +340,8 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
|
||||
# --- Case 2: 3D (fast cross-product method) ---
|
||||
elif n_dims * nd_fixup == 3:
|
||||
num_random_fields = 2
|
||||
|
||||
kn_z, kn_y, kn_x = k_norm_components
|
||||
ez = torch.tensor([0.0, 0.0, 1.0], device=self.device, dtype=self.dtype)
|
||||
is_parallel = (kn_x.abs() < 1e-6) & (kn_y.abs() < 1e-6)
|
||||
@@ -441,7 +363,6 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
v = (vz, vy, vx)
|
||||
|
||||
basis_vectors = [u, v]
|
||||
num_random_fields = 2
|
||||
|
||||
# --- Case 3: N-D (General Gram-Schmidt process) ---
|
||||
else:
|
||||
@@ -505,7 +426,7 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
|
||||
_handle_field_basis_ndim = _handle_field_basis
|
||||
|
||||
def generate_octaves(
|
||||
def generate_field(
|
||||
self,
|
||||
batch: int,
|
||||
height: int,
|
||||
@@ -567,7 +488,7 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
def generate(self, *args) -> torch.Tensor:
|
||||
cm = self.channel_mode
|
||||
if self.noise_chunk is None:
|
||||
self.noise_chunk = self.generate_octaves(
|
||||
self.noise_chunk = self.generate_field(
|
||||
self.eff_batch,
|
||||
self.height,
|
||||
self.width,
|
||||
@@ -585,6 +506,9 @@ class SimulationNoiseGenerator(NoiseGenerator):
|
||||
),
|
||||
dim=2,
|
||||
)
|
||||
elif cm == "flat":
|
||||
noise = self.noise_chunk[:, :, self.current_depth]
|
||||
noise = noise.flatten()[: math.prod(self.shape)]
|
||||
elif cm in {"over_depth", "over_depth_alt"}:
|
||||
noise = self.noise_chunk[:, :, depth_from:depth_to]
|
||||
if cm == "over_depth":
|
||||
|
||||
+324
-5
@@ -2,8 +2,9 @@ from __future__ import annotations
|
||||
|
||||
import math
|
||||
import random
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Callable
|
||||
from enum import Enum, auto
|
||||
from functools import lru_cache, partial
|
||||
from typing import TYPE_CHECKING, NamedTuple
|
||||
|
||||
import torch
|
||||
from comfy.model_management import device_supports_non_blocking, get_torch_device
|
||||
@@ -12,12 +13,15 @@ from comfy.utils import common_upscale
|
||||
from .external import MODULES as EXT
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Sequence
|
||||
from collections.abc import Callable, Sequence
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
BLENDING_MODES = {
|
||||
"lerp": torch.lerp,
|
||||
"inject": lambda a, b, t: (b * t).add_(a),
|
||||
"subtract_b": lambda a, b, t: a - b * t,
|
||||
"weighted_average": lambda a, b, t: (b * t).add_(a) / (1.0 + abs(t)),
|
||||
}
|
||||
UPSCALE_METHODS = (
|
||||
"bilinear",
|
||||
@@ -95,7 +99,7 @@ def scale_noise(
|
||||
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
|
||||
noise = (noise / std).nan_to_num_()
|
||||
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)
|
||||
@@ -121,6 +125,47 @@ def tensor_to(
|
||||
return tensor.to(dest, non_blocking=non_blocking)
|
||||
|
||||
|
||||
def range_wrap(
|
||||
x: torch.Tensor,
|
||||
min_val: float | torch.Tensor,
|
||||
max_val: float | torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return min_val + (x - min_val).remainder_(max_val - min_val)
|
||||
|
||||
|
||||
def softplus_soft_clamp(
|
||||
t: torch.Tensor,
|
||||
min_val: torch.Tensor | float = 0.0,
|
||||
max_val: torch.Tensor | float = 1.0,
|
||||
*,
|
||||
# We define stiffness as a multiplier (beta) for the softplus function.
|
||||
# Higher stiffness = sharper transition.
|
||||
stiffness: float = 1.0,
|
||||
safe: bool = True,
|
||||
) -> torch.Tensor:
|
||||
if isinstance(min_val, (float, int)):
|
||||
min_val = t.new_tensor(min_val)
|
||||
if isinstance(max_val, (float, int)):
|
||||
max_val = t.new_tensor(max_val)
|
||||
|
||||
if stiffness < 1e-04:
|
||||
return t.clamp(min_val, max_val)
|
||||
|
||||
# Calculate how much we are exceeding the Max
|
||||
# softplus(beta * x) / beta
|
||||
upper_overshoot = F.softplus((t - max_val).mul_(stiffness)).div_(-stiffness)
|
||||
|
||||
# Calculate how much we are falling short of the Min
|
||||
lower_undershoot = F.softplus((min_val - t).mul_(stiffness)).div_(stiffness)
|
||||
|
||||
# Apply corrections:
|
||||
# Original - (Amount over max) + (Amount under min)
|
||||
t = upper_overshoot.add_(t).add_(lower_undershoot)
|
||||
if safe:
|
||||
t = t.clamp(min_val, max_val)
|
||||
return t
|
||||
|
||||
|
||||
def _quantile_norm_scaledown(
|
||||
noise: torch.Tensor,
|
||||
nq: torch.Tensor,
|
||||
@@ -360,10 +405,51 @@ quantile_handlers = {
|
||||
count_flipping=True,
|
||||
avoid_sign=True,
|
||||
),
|
||||
"wrap": lambda noise, nq, **_kwargs: range_wrap(noise, -nq, nq),
|
||||
"wrap_keepsign": lambda noise, nq, **_kwargs: torch.where(
|
||||
noise.abs() > nq,
|
||||
range_wrap(noise, -nq, nq).copysign_(noise),
|
||||
noise,
|
||||
),
|
||||
"wrap_avoidsign": lambda noise, nq, **_kwargs: torch.where(
|
||||
noise.abs() > nq,
|
||||
range_wrap(noise, -nq, nq).copysign_(noise.neg()),
|
||||
noise,
|
||||
),
|
||||
"softplus_clamp_s01": lambda noise, nq, **_kwargs: softplus_soft_clamp(
|
||||
noise,
|
||||
-nq,
|
||||
nq,
|
||||
stiffness=0.1,
|
||||
),
|
||||
"softplus_clamp_s05": lambda noise, nq, **_kwargs: softplus_soft_clamp(
|
||||
noise,
|
||||
-nq,
|
||||
nq,
|
||||
stiffness=0.5,
|
||||
),
|
||||
"softplus_clamp_s1": lambda noise, nq, **_kwargs: softplus_soft_clamp(
|
||||
noise,
|
||||
-nq,
|
||||
nq,
|
||||
stiffness=1.0,
|
||||
),
|
||||
"softplus_clamp_s2": lambda noise, nq, **_kwargs: softplus_soft_clamp(
|
||||
noise,
|
||||
-nq,
|
||||
nq,
|
||||
stiffness=1.0,
|
||||
),
|
||||
"softplus_clamp_s5": lambda noise, nq, **_kwargs: softplus_soft_clamp(
|
||||
noise,
|
||||
-nq,
|
||||
nq,
|
||||
stiffness=5.0,
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
# Initial version based on Studentt distribution normalizatino from https://github.com/Clybius/ComfyUI-Extra-Samplers/
|
||||
# Initial version based on StudentT distribution normalization from https://github.com/Clybius/ComfyUI-Extra-Samplers/
|
||||
def quantile_normalize(
|
||||
noise: torch.Tensor,
|
||||
*,
|
||||
@@ -449,6 +535,239 @@ def quantile_normalize(
|
||||
return noise if noise.shape == orig_shape else noise.reshape(orig_shape)
|
||||
|
||||
|
||||
# class QuantileNormMode(Enum):
|
||||
# # Quantile is applied to absolute values.
|
||||
# SYMMETRIC = auto()
|
||||
# # Quantile is applied to signed values.
|
||||
# SEPERATE = auto()
|
||||
|
||||
|
||||
class QuantileNormQuantileMode(Enum):
|
||||
QUANTILE = auto()
|
||||
# User supplied value to use as the quantile value.
|
||||
USER_HIGH = auto()
|
||||
# Like setting a negative quantile.
|
||||
USER_LOW = auto()
|
||||
|
||||
|
||||
class QuantileNormSignMode(Enum):
|
||||
DEFAULT = auto()
|
||||
KEEP = auto()
|
||||
AVOID = auto()
|
||||
|
||||
|
||||
class QuantileNormTargetMode(Enum):
|
||||
BOTH = auto()
|
||||
POSITIVE = auto()
|
||||
NEGATIVE = auto()
|
||||
|
||||
|
||||
class QuantileNorm(NamedTuple):
|
||||
# mode: QuantileNormMode = QuantileNormMode.SYMMETRIC
|
||||
quantile_mode: QuantileNormQuantileMode = QuantileNormQuantileMode.QUANTILE
|
||||
sign_mode: QuantileNormSignMode = QuantileNormSignMode.DEFAULT
|
||||
target_mode: QuantileNormTargetMode = QuantileNormTargetMode.BOTH
|
||||
strategy: str = "clamp"
|
||||
# Overrides strategy.
|
||||
strategy_handler: Callable | None = None
|
||||
start_end_dim: tuple[int, int] | None = (1, -1)
|
||||
dims: tuple[int, ...] = ()
|
||||
quantile: float | torch.Tensor = 0.75
|
||||
# If none, we use absmax when quantile is negative, otherwise
|
||||
# abs value at this quantile.
|
||||
low_extreme_quantile: float | None = 0.99
|
||||
nq_scale: float = 1.0
|
||||
power: float = 0.0
|
||||
use_abs_quantile: bool = True
|
||||
# When setting a target other than BOTH, apply the mask to the quantile calculation as well.
|
||||
use_quantile_mask: bool = True
|
||||
use_float64: bool = True
|
||||
fix_invalid: bool = True
|
||||
|
||||
@staticmethod
|
||||
def fix_dim(dim: int, ndim: int) -> int | None:
|
||||
if dim < 0:
|
||||
dim = ndim + dim
|
||||
return dim if 0 <= dim < ndim else None
|
||||
|
||||
@classmethod
|
||||
def fix_dims(cls, dims: tuple[int, ...], ndim: int) -> tuple[int, ...]:
|
||||
dims = tuple(d for d in (cls.fix_dim(d_, ndim) for d_ in dims) if d is not None)
|
||||
return tuple({dims})
|
||||
|
||||
def get_dims(self, ndim: int) -> tuple[int, ...]:
|
||||
dims = self.fix_dims(self.dims)
|
||||
if self.start_end_dim is None:
|
||||
return dims
|
||||
start_end_dim = self.fix_dims(*self.start_end_dim, ndim)
|
||||
if len(start_end_dim) != 2:
|
||||
return dims
|
||||
sd, ed = start_end_dim
|
||||
if sd > ed:
|
||||
sd, ed = ed, sd
|
||||
return tuple({range(sd, ed + 1), *dims})
|
||||
|
||||
@classmethod
|
||||
@lru_cache(maxsize=128)
|
||||
def get_perms(
|
||||
cls,
|
||||
# Must be deduped and sanitized.
|
||||
dims: tuple[int, ...],
|
||||
ndim: int,
|
||||
) -> tuple[tuple[int, ...], tuple[int, ...]]:
|
||||
other_dims = tuple(d for d in range(ndim) if d not in dims)
|
||||
perms = (*other_dims, *dims)
|
||||
inv_perms_t = torch.nn.utils.rnn.invert_permutation(
|
||||
torch.tensor(perms, device="cpu"),
|
||||
)
|
||||
if inv_perms_t is None:
|
||||
errstr = f"torch.nn.utils.rnn.invert_permutation returned None for input {perms}!"
|
||||
raise RuntimeError(errstr)
|
||||
inv_perms = tuple(inv_perms_t.tolist())
|
||||
return (perms, inv_perms)
|
||||
|
||||
def __call__(self, t: torch.Tensor) -> torch.Tensor:
|
||||
ndim = t.ndim
|
||||
dims = self.get_dims(ndim)
|
||||
handler = (
|
||||
quantile_handlers.get(self.strategy)
|
||||
if self.strategy_handler is None
|
||||
else self.strategy_handler
|
||||
)
|
||||
if handler is None:
|
||||
raise ValueError("No strategy handler")
|
||||
if not dims or self.quantile == 0 or t.numel() < 2:
|
||||
return t
|
||||
orig_dtype = t.dtype
|
||||
orig_t = t
|
||||
eff_dtype = torch.float64 if self.use_float64 else torch.float32
|
||||
dlen = len(dims)
|
||||
olen = ndim - dlen
|
||||
perms, inv_perms = self.get_perms(dims, ndim)
|
||||
t = t.to(dtype=eff_dtype)
|
||||
# Move the dims we're working with to the end.
|
||||
t = t.permute(perms)
|
||||
permuted_shape = t.shape
|
||||
# And flatten them.
|
||||
t = t.flatten(start_dim=olen)
|
||||
use_low_extreme = (
|
||||
isinstance(self.quantile, float) and self.quantile < 0
|
||||
) or self.quantile_mode == QuantileNormQuantileMode.USER_LOW
|
||||
if use_low_extreme:
|
||||
raise RuntimeError("NYI")
|
||||
if self.target_mode == QuantileNormTargetMode.NEGATIVE:
|
||||
mask = t.sign() < 0
|
||||
elif self.target_mode == QuantileNormTargetMode.POSITIVE:
|
||||
mask = t.sign() > 0
|
||||
else:
|
||||
mask = None
|
||||
masked_quantile = self.use_quantile_mask and mask is not None
|
||||
if self.quantile_mode == QuantileNormQuantileMode.QUANTILE:
|
||||
nq_input = t.abs() if self.use_abs_quantile else t
|
||||
if masked_quantile:
|
||||
nq_input = nq_input[mask]
|
||||
if not torch.any(nq_input):
|
||||
return orig_t
|
||||
nq = torch.quantile(nq_input, dim=-1, keepdim=True)
|
||||
else:
|
||||
raise RuntimeError("NYI")
|
||||
if self.nq_scale != 1.0:
|
||||
nq *= self.nq_scale
|
||||
# ...
|
||||
if self.power not in {0.0, 1.0}:
|
||||
t = t.abs().pow_(self.power).copysign_(t)
|
||||
if self.fix_invalid:
|
||||
t = t.nan_to_num_()
|
||||
t = t.to(dtype=orig_dtype)
|
||||
# Back to the original shape.
|
||||
return t.reshape(permuted_shape).permute(inv_perms).contiguous()
|
||||
|
||||
|
||||
def quantile_normalize_adv(
|
||||
noise: torch.Tensor,
|
||||
*,
|
||||
quantile: float | tuple | list = 0.75,
|
||||
dim: int | None = 1,
|
||||
flatten: bool = True,
|
||||
nq_fac: float = 1.0,
|
||||
pow_fac: float = 0.5,
|
||||
strategy: str = "clamp",
|
||||
strategy_handler=None,
|
||||
eps=1e-08,
|
||||
) -> torch.Tensor:
|
||||
if noise.numel() == 0:
|
||||
return noise
|
||||
if isinstance(quantile, (tuple, list)):
|
||||
for q in quantile:
|
||||
noise = quantile_normalize(
|
||||
noise=noise,
|
||||
quantile=q,
|
||||
dim=dim,
|
||||
flatten=flatten,
|
||||
nq_fac=nq_fac,
|
||||
pow_fac=pow_fac,
|
||||
strategy=strategy,
|
||||
strategy_handler=strategy_handler,
|
||||
)
|
||||
return noise
|
||||
if quantile is None or quantile >= 1 or quantile <= -1:
|
||||
return noise
|
||||
centered = quantile < 0
|
||||
absquantile = abs(quantile)
|
||||
orig_shape = noise.shape
|
||||
if noise.ndim > 1 and flatten:
|
||||
flatnoise = noise.flatten(start_dim=dim)
|
||||
else:
|
||||
flatten = False
|
||||
flatnoise = noise
|
||||
handler = (
|
||||
quantile_handlers.get(strategy)
|
||||
if strategy_handler is None
|
||||
else strategy_handler
|
||||
)
|
||||
if handler is None:
|
||||
raise ValueError("Unknown strategy")
|
||||
if not centered:
|
||||
nq = torch.quantile(
|
||||
flatnoise.abs(),
|
||||
quantile,
|
||||
dim=-1 if flatten else dim,
|
||||
keepdim=True,
|
||||
)
|
||||
nq = nq.mul_(nq_fac).add_(eps)
|
||||
# print(f"\nNQ: {nq}")
|
||||
noise = handler(
|
||||
flatnoise,
|
||||
nq,
|
||||
orig_noise=noise,
|
||||
dim=dim,
|
||||
flatten=flatten,
|
||||
)
|
||||
else:
|
||||
absnoise = flatnoise.abs()
|
||||
maxabs = absnoise.amax(dim=-1 if flatten else dim, keepdim=True)
|
||||
proxy = flatnoise.sign().mul_(maxabs - absnoise)
|
||||
nq_proxy = torch.quantile(
|
||||
proxy.abs(),
|
||||
absquantile,
|
||||
dim=-1 if flatten else dim,
|
||||
keepdim=True,
|
||||
)
|
||||
nq_proxy = nq_proxy.mul_(nq_fac).add_(eps)
|
||||
# print(f"\nNQ proxy: {nq_proxy}")
|
||||
out_proxy = handler(
|
||||
proxy,
|
||||
nq_proxy,
|
||||
orig_noise=noise,
|
||||
dim=dim,
|
||||
flatten=flatten,
|
||||
)
|
||||
noise = out_proxy.sign().mul_(maxabs - out_proxy.abs())
|
||||
if pow_fac not in {0.0, 1.0}:
|
||||
noise = noise.abs().pow_(pow_fac).copysign(noise)
|
||||
return noise if noise.shape == orig_shape else noise.reshape(orig_shape)
|
||||
|
||||
|
||||
def normalize_to_scale(
|
||||
latent: torch.Tensor,
|
||||
target_min: float,
|
||||
|
||||
+5
-1
@@ -172,6 +172,7 @@ class WCFGPercentages(NamedTuple):
|
||||
else:
|
||||
pct_enabled_sigmas = (start_sigma - sigma) / (start_sigma - end_sigma)
|
||||
steps = len(sigmas) - 1
|
||||
have_steps = False
|
||||
if steps > 1:
|
||||
step = utils.step_from_sigmas(sigma, sigmas)
|
||||
pct_steps = step / (steps - 1) if step is not None else None
|
||||
@@ -179,10 +180,11 @@ class WCFGPercentages(NamedTuple):
|
||||
(sigmas <= start_sigma) & (sigmas >= end_sigma)
|
||||
]
|
||||
if len(enabled_steps) > 1:
|
||||
have_steps = True
|
||||
step_first = enabled_steps[0].item()
|
||||
step_last = enabled_steps[-1].item()
|
||||
pct_enabled_steps = (step - step_first) / (step_last - step_first)
|
||||
else:
|
||||
if not have_steps:
|
||||
step = 0.0
|
||||
pct_steps = 1.0
|
||||
step_first = step_last = None
|
||||
@@ -247,6 +249,8 @@ class WCFGScales(NamedTuple):
|
||||
target = self.yh_scales
|
||||
if isinstance(target, float):
|
||||
return f"{target:.4f}"
|
||||
if not isinstance(target, (list, tuple)):
|
||||
return str(target)
|
||||
result = ", ".join(
|
||||
self.pretty_yh_scales(target=val)
|
||||
if isinstance(val, (list, tuple))
|
||||
|
||||
Reference in New Issue
Block a user