Sync changes

This commit is contained in:
blepping
2026-03-02 02:50:41 -07:00
parent 4ec5970128
commit ec7def5723
12 changed files with 1522 additions and 835 deletions
+7 -1
View File
@@ -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):
+12 -5
View File
@@ -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
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+4 -3
View File
@@ -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
View File
@@ -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
+2
View File
@@ -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
+33 -109
View File
@@ -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
View File
@@ -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
View File
@@ -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))