First commit.

Add RES_Momentumized
Add DPMPP_DualSDE_Momentumized
add egotistical Clyb_4M_SDE_Momentumized
add TTM (Doesn't seem to work well with current imp.?)
add LCM Custom Noise
This commit is contained in:
Clybius
2024-02-02 09:15:29 -06:00
commit 9f08cd98a4
5 changed files with 1594 additions and 0 deletions
+17
View File
@@ -0,0 +1,17 @@
from . import extra_samplers
from . import nodes
extra_samplers.add_samplers()
#extra_samplers.add_schedulers()
NODE_CLASS_MAPPINGS = {
"SamplerCustomNoise": nodes.SamplerCustomNoise,
"SamplerCustomNoiseDuo": nodes.SamplerCustomNoiseDuo,
"SamplerCustomModelMixtureDuo": nodes.SamplerCustomModelMixtureDuo,
"SamplerRES_Momentumized": nodes.SamplerRES_MOMENTUMIZED,
"SamplerDPMPP_DualSDE_Momentumized": nodes.SamplerDPMPP_DUALSDE_MOMENTUMIZED,
"SamplerCLYB_4M_SDE_Momentumized": nodes.SamplerCLYB_4M_SDE_MOMENTUMIZED,
"SamplerTTM": nodes.SamplerTTM,
"SamplerLCMCustom": nodes.SamplerLCMCustom,
}
__all__ = ['NODE_CLASS_MAPPINGS']
+689
View File
@@ -0,0 +1,689 @@
import math
from scipy import integrate
import torch
from torch import nn
import torchsde
from tqdm.auto import trange, tqdm
import comfy.sample
import k_diffusion.sampling
from k_diffusion.sampling import BrownianTreeNoiseSampler, PIDStepSizeController, get_ancestral_step, to_d, default_noise_sampler
# The following function adds the samplers during initialization, in __init__.py
def add_samplers():
from comfy.samplers import KSampler, k_diffusion_sampling
for sampler in extra_samplers: #getattr(self, "sample_{}".format(extra_samplers))
if sampler not in KSampler.SAMPLERS:
try:
idx = KSampler.SAMPLERS.index("uni_pc_bh2") # Last item in the samplers list
KSampler.SAMPLERS.insert(idx+1, sampler) # Add our custom samplers
setattr(k_diffusion_sampling, "sample_{}".format(sampler), extra_samplers[sampler])
import importlib
importlib.reload(k_diffusion_sampling)
except ValueError as err:
pass
# The following function adds the samplers during initialization, in __init__.py
def add_schedulers():
from comfy.samplers import KSampler, k_diffusion_sampling
for scheduler in extra_schedulers: #getattr(self, "sample_{}".format(extra_samplers))
if scheduler not in KSampler.SCHEDULERS:
try:
idx = KSampler.SCHEDULERS.index("ddim_uniform") # Last item in the samplers list
KSampler.SCHEDULERS.insert(idx+1, scheduler) # Add our custom samplers
setattr(k_diffusion_sampling, "get_sigmas_{}".format(scheduler), extra_schedulers[scheduler])
import importlib
importlib.reload(k_diffusion_sampling)
except ValueError as err:
pass
# Noise samplers
from torch import Generator, Tensor, lerp
from torch.nn.functional import unfold
from typing import Callable, Tuple
from math import pi
def get_positions(block_shape: Tuple[int, int]) -> Tensor:
"""
Generate position tensor.
Arguments:
block_shape -- (height, width) of position tensor
Returns:
position vector shaped (1, height, width, 1, 1, 2)
"""
bh, bw = block_shape
positions = torch.stack(
torch.meshgrid(
[(torch.arange(b) + 0.5) / b for b in (bw, bh)],
indexing="xy",
),
-1,
).view(1, bh, bw, 1, 1, 2)
return positions
def unfold_grid(vectors: Tensor) -> Tensor:
"""
Unfold vector grid to batched vectors.
Arguments:
vectors -- grid vectors
Returns:
batched grid vectors
"""
batch_size, _, gpy, gpx = vectors.shape
return (
unfold(vectors, (2, 2))
.view(batch_size, 2, 4, -1)
.permute(0, 2, 3, 1)
.view(batch_size, 4, gpy - 1, gpx - 1, 2)
)
def smooth_step(t: Tensor) -> Tensor:
"""
Smooth step function [0, 1] -> [0, 1].
Arguments:
t -- input values (any shape)
Returns:
output values (same shape as input values)
"""
return t * t * (3.0 - 2.0 * t)
def perlin_noise_tensor(
vectors: Tensor, positions: Tensor, step: Callable = None
) -> Tensor:
"""
Generate perlin noise from batched vectors and positions.
Arguments:
vectors -- batched grid vectors shaped (batch_size, 4, grid_height, grid_width, 2)
positions -- batched grid positions shaped (batch_size or 1, block_height, block_width, grid_height or 1, grid_width or 1, 2)
Keyword Arguments:
step -- smooth step function [0, 1] -> [0, 1] (default: `smooth_step`)
Raises:
Exception: if position and vector shapes do not match
Returns:
(batch_size, block_height * grid_height, block_width * grid_width)
"""
if step is None:
step = smooth_step
batch_size = vectors.shape[0]
# grid height, grid width
gh, gw = vectors.shape[2:4]
# block height, block width
bh, bw = positions.shape[1:3]
for i in range(2):
if positions.shape[i + 3] not in (1, vectors.shape[i + 2]):
raise Exception(
f"Blocks shapes do not match: vectors ({vectors.shape[1]}, {vectors.shape[2]}), positions {gh}, {gw})"
)
if positions.shape[0] not in (1, batch_size):
raise Exception(
f"Batch sizes do not match: vectors ({vectors.shape[0]}), positions ({positions.shape[0]})"
)
vectors = vectors.view(batch_size, 4, 1, gh * gw, 2)
positions = positions.view(positions.shape[0], bh * bw, -1, 2)
step_x = step(positions[..., 0])
step_y = step(positions[..., 1])
row0 = lerp(
(vectors[:, 0] * positions).sum(dim=-1),
(vectors[:, 1] * (positions - positions.new_tensor((1, 0)))).sum(dim=-1),
step_x,
)
row1 = lerp(
(vectors[:, 2] * (positions - positions.new_tensor((0, 1)))).sum(dim=-1),
(vectors[:, 3] * (positions - positions.new_tensor((1, 1)))).sum(dim=-1),
step_x,
)
noise = lerp(row0, row1, step_y)
return (
noise.view(
batch_size,
bh,
bw,
gh,
gw,
)
.permute(0, 3, 1, 4, 2)
.reshape(batch_size, gh * bh, gw * bw)
)
def perlin_noise(
grid_shape: Tuple[int, int],
out_shape: Tuple[int, int],
batch_size: int = 1,
generator: Generator = None,
*args,
**kwargs,
) -> Tensor:
"""
Generate perlin noise with given shape. `*args` and `**kwargs` are forwarded to `Tensor` creation.
Arguments:
grid_shape -- Shape of grid (height, width).
out_shape -- Shape of output noise image (height, width).
Keyword Arguments:
batch_size -- (default: {1})
generator -- random generator used for grid vectors (default: {None})
Raises:
Exception: if grid and out shapes do not match
Returns:
Noise image shaped (batch_size, height, width)
"""
# grid height and width
gh, gw = grid_shape
# output height and width
oh, ow = out_shape
# block height and width
bh, bw = oh // gh, ow // gw
if oh != bh * gh:
raise Exception(f"Output height {oh} must be divisible by grid height {gh}")
if ow != bw * gw != 0:
raise Exception(f"Output width {ow} must be divisible by grid width {gw}")
angle = torch.empty(
[batch_size] + [s + 1 for s in grid_shape], *args, **kwargs
).uniform_(to=2.0 * pi, generator=generator)
# random vectors on grid points
vectors = unfold_grid(torch.stack((torch.cos(angle), torch.sin(angle)), dim=1))
# positions inside grid cells [0, 1)
positions = get_positions((bh, bw)).to(vectors)
return perlin_noise_tensor(vectors, positions).squeeze(0)
def rand_perlin_like(x):
noise = torch.randn_like(x) / 2.0
noise_size_H = noise.size(dim=2)
noise_size_W = noise.size(dim=3)
perlin = None
for i in range(2):
noise += perlin_noise((noise_size_H, noise_size_W), (noise_size_H, noise_size_W), batch_size=4).to(x.device)
#noise += perlin
#print(noise)
return noise / noise.std()
def uniform_noise_sampler(x): # Even distribution, seemingly produces more information in non-subject areas than the normal (gaussian) noise sampler
return lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
from torch.distributions import StudentT
def studentt_noise_sampler(x): # Produces more subject-focused outputs due to distribution, unsure if this works
noise = StudentT(loc=0, scale=0.2, df=1).rsample(x.size())
#noise *= 2 / (torch.max(torch.abs(noise)) + 1e-8)
s: FloatTensor = torch.quantile(
noise.flatten(start_dim=1).abs(),
0.75,
dim = -1
)
#s.clamp_(min = 1.)
s = s.reshape(*s.shape, 1, 1, 1)
noise = noise.clamp(-s, s)
noise = torch.copysign(torch.pow(torch.abs(noise), 0.5), noise)
print(s)
return lambda sigma, sigma_next: noise.to(x.device) / (7/3)
def highres_pyramid_noise_like(x, discount=0.7):
b, c, h, w = x.shape # EDIT: w and h get over-written, rename for a different variant!
orig_h = h
orig_w = w
u = torch.nn.Upsample(size=(orig_h, orig_w), mode='bilinear')
noise = (torch.rand_like(x) - 0.5) * 2 * 1.73 # Start with scaled uniform noise
for i in range(4):
r = random.random()*2+2 # Rather than always going 2x,
h, w = min(orig_h*15, int(h*(r**i))), min(orig_w*15, int(w*(r**i)))
noise += u(torch.randn(b, c, h, w).to(x)) * discount**i
if h>=orig_h*15 or w>=orig_w*15: break # Lowest resolution is 1x1
return noise/noise.std() # Scaled back to roughly unit variance
def green_noise_like(x):
noise = torch.randn_like(x)
width = noise.size(dim=2)
height = noise.size(dim=3)
scale = 1.0 / (width * height)
fy = torch.fft.fftfreq(width, device=x.device)[:, None] ** 2
fx = torch.fft.fftfreq(height, device=x.device) ** 2
f = fy + fx
power = torch.sqrt(f)
power[0, 0] = 1
noise = torch.fft.ifft2(torch.fft.fft2(noise) / torch.sqrt(power))
noise *= scale / noise.std()
noise = torch.real(noise).to(x.device)
return noise / noise.std()
def green_noise_sampler(x): # This doesn't work properly right now
width = x.size(dim=2)
height = x.size(dim=3)
noise = torch.randn(width, height)
#scale = 1.0 / (width * height)
fy = torch.fft.fftfreq(width)[:, None] ** 2
fx = torch.fft.fftfreq(height) ** 2
f = fy + fx
power = torch.sqrt(f)
power[0, 0] = 1
noise = torch.fft.ifft2(torch.fft.fft2(noise) / torch.sqrt(power))
#noise *= scale / noise.std()
noise = torch.real(noise).to(x.device)
mean = torch.mean(noise)
std = torch.std(noise)
noise.sub_(mean).div_(std)
print(noise)
return lambda sigma, sigma_next: noise
def power_noise_sampler(tensor, alpha=2, k=1): # This doesn't work properly right now
"""Generate 1/f noise for a given tensor.
Args:
tensor: The tensor to add noise to.
alpha: The parameter that determines the slope of the spectrum.
k: A constant.
Returns:
A tensor with the same shape as `tensor` containing 1/f noise.
"""
tensor = torch.randn_like(tensor)
fft = torch.fft.fft2(tensor)
freq = torch.arange(1, len(fft) + 1, dtype=torch.float)
spectral_density = k / freq**alpha
noise = torch.rand(tensor.shape) * spectral_density
mean = torch.mean(noise, dim=(-2, -1), keepdim=True).to(tensor.device)
std = torch.std(noise, dim=(-2, -1), keepdim=True).to(tensor.device)
noise = noise.to(tensor.device).sub_(mean).div_(std)
variance = torch.var(noise, dim=(-2, -1), keepdim=True)
print(variance)
return lambda sigma, sigma_next: noise / 3
def mixed_noise_sampler(x): # Meant for RES
gaussian = torch.randn_like(x)
uniform = ((torch.rand_like(x) - 0.5) * 2 * 1.73)
# Calculate variances
#gaussian_variance = torch.var(gaussian, dim=(-2, -1), keepdim=True)
#uniform_variance = torch.var(uniform, dim=(-2, -1), keepdim=True)
# Determine weights based on variances
#total_variance = gaussian_variance + uniform_variance
#gaussian_weight = gaussian_variance / total_variance
#uniform_weight = uniform_variance / total_variance
#mixed_noise = gaussian * gaussian_weight + uniform * uniform_weight
#print(gaussian_weight, uniform_weight)
mixed_noise = (gaussian * uniform) / 4
# Return the final mixed noise sample
return lambda sigma, sigma_next: mixed_noise
import random
def pyramid_noise_like(x, discount=0.75):
b, c, w, h = x.shape # EDIT: w and h get over-written, rename for a different variant!
#noise_vector_magnitude = (torch.linalg.vector_norm(torch.randn_like(x), dim=(1)) + 0.0000000001)[:,None]
gauss_noise = torch.randn_like(x)
gn_mean = gauss_noise.mean() * 0
gn_std = gauss_noise.std() / 2
noise = torch.nn.functional.interpolate((torch.normal(mean=0, std=0.25, size=(b, c, w // 2, h // 2)).to(x)), size=(w, h), mode='bicubic')
noise_2 = torch.nn.functional.interpolate((torch.normal(mean=0, std=0.5, size=(b, c, w * 8, h * 8)).to(x)), size=(w, h), mode='bicubic')
noise = noise + noise_2
#gauss_mag = (torch.linalg.vector_norm(gauss_noise, dim=(1)) + 0.0000000001)[:,None]
#noise_mag = (torch.linalg.vector_norm(noise, dim=(1)) + 0.0000000001)[:,None]
#noise /= noise_mag
#noise *= gauss_mag
#noise_scaled = torch.copysign(torch.pow(torch.abs(noise / noise.max()), 0.95), noise) * 1.4
#noise = torch.copysign(torch.pow(torch.abs(noise_scaled), 0.9), noise_scaled) / 1.15
return lambda sigma, sigma_next: noise# / 2# / 1.5#torch.copysign(torch.pow(torch.abs(noise), 1.25), noise) * 2 * 1.73#.sub_(noise.mean()).div_(noise.std()) # Scaled back to roughly unit variance
# Below this point are extra samplers
@torch.no_grad()
def sample_clyb_4m_sde_momentumized(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1.0, s_noise=1., noise_sampler=None, momentum=0.5):
"""DPM-Solver++(3M) SDE, modified with an extra SDE, and momentumized in both the SDE and ODE(?). 'its a first' - Clybius 2023
The expression for d1 is derived from the extrapolation formula given in the paper “Diffusion Monte Carlo with stochastic Hamiltonians” by M. Foulkes, L. Mitas, R. Needs, and G. Rajagopal. The formula is given as follows:
d1 = d1_0 + (d1_0 - d1_1) * r2 / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) * r2 / ((r2 + r1) * (r0 + r1))
(if this is an incorrect citing, we blame Google's Bard and OpenAI's ChatGPT for this and NOT me :^) )
where d1_0, d1_1, and d1_2 are defined as follows:
d1_0 = (denoised - denoised_1) / r2
d1_1 = (denoised_1 - denoised_2) / r1
d1_2 = (denoised_2 - denoised_3) / r0
The variables r0, r1, and r2 are defined as follows:
r0 = h_3 / h_2
r1 = h_2 / h
r2 = h / h_1
"""
def momentum_func(diff, velocity, timescale=1.0, offset=-momentum / 2.0): # Diff is current diff, vel is previous diff
if velocity is None:
momentum_vel = diff
else:
momentum_vel = momentum * (timescale + offset) * velocity + (1 - momentum * (timescale + offset)) * diff
return momentum_vel
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
noise_sampler = rand_perlin_like(x) if noise_sampler is None else noise_sampler
extra_args = {} if extra_args is None else extra_args
s_in = x.new_ones([x.shape[0]])
denoised_1, denoised_2, denoised_3 = None, None, None
h_1, h_2, h_3 = None, None, None
vel, vel_sde = None, None
for i in trange(len(sigmas) - 1, disable=disable):
time = sigmas[i] / sigma_max
denoised = model(x, sigmas[i] * s_in, **extra_args)
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
if sigmas[i + 1] == 0:
# Denoising step
x = denoised
else:
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
h = s - t
h_eta = h * (eta + 1)
x_diff = momentum_func((-h_eta).expm1().neg() * denoised, vel, time)
vel = x_diff
x = torch.exp(-h_eta) * x + vel
if h_3 is not None:
r0 = h_3 / h_2
r1 = h_2 / h
r2 = h / h_1
d1_0 = (denoised - denoised_1) / r2
d1_1 = (denoised_1 - denoised_2) / r1
d1_2 = (denoised_2 - denoised_3) / r0
d1 = d1_0 + (d1_0 - d1_1) * r2 / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) * r2 / ((r2 + r1) * (r0 + r1))
d2 = (d1_0 - d1_1) / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) / ((r2 + r1) * (r0 + r1))
phi_3 = h_eta.neg().expm1() / h_eta + 1
phi_4 = phi_3 / h_eta - 0.5
sde_diff = momentum_func(phi_3 * d1 - phi_4 * d2, vel_sde, time)
vel_sde = sde_diff
x = x + vel_sde
elif h_2 is not None:
r0 = h_1 / h
r1 = h_2 / h
d1_0 = (denoised - denoised_1) / r0
d1_1 = (denoised_1 - denoised_2) / r1
d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1)
d2 = (d1_0 - d1_1) / (r0 + r1)
phi_2 = h_eta.neg().expm1() / h_eta + 1
phi_3 = phi_2 / h_eta - 0.5
sde_diff = momentum_func(phi_2 * d1 - phi_3 * d2, vel_sde, time)
vel_sde = sde_diff
x = x + vel_sde
elif h_1 is not None:
r = h_1 / h
d = (denoised - denoised_1) / r
phi_2 = h_eta.neg().expm1() / h_eta + 1
sde_diff = momentum_func(phi_2 * d, vel_sde, time)
vel_sde = sde_diff
x = x + vel_sde
if eta:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * h * eta).expm1().neg().sqrt() * s_noise
denoised_1, denoised_2, denoised_3 = denoised, denoised_1, denoised_2
h_1, h_2, h_3 = h, h_1, h_2
return x
# Kat's Truncated Taylor Method sampler, by Katherine Crowson
def sample_ttm_jvp(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
"""Second order truncated Taylor method (torch.func.jvp() version)."""
extra_args = {} if extra_args is None else extra_args
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
s_in = x.new_ones([x.shape[0]])
model_fn = lambda x, sigma: model(x, sigma * s_in, **extra_args)
for i in trange(len(sigmas) - 1, disable=disable):
denoised = model_fn(x, sigmas[i])
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
if sigmas[i + 1] == 0:
# Denoising step
x = denoised
else:
# 2nd order truncated Taylor method
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
h = s - t
h_eta = h * (eta + 1)
eps = to_d(x, sigmas[i], denoised)
_, denoised_prime = torch.func.jvp(model_fn, (x, sigmas[i]), (eps * -sigmas[i], -sigmas[i]))
phi_1 = -torch.expm1(-h_eta)
#phi_2 = torch.expm1(-h_eta) + h_eta
phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0
x = torch.exp(-h_eta) * x + phi_1 * denoised + phi_2 * denoised_prime
if eta:
phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * eta))
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * phi_1_noise * s_noise
return x
from .other_samplers.refined_exp_solver import sample_refined_exp_s
def sample_res_solver(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None, denoise_to_zero=True, simple_phi_calc=False, c2=0.5, ita=torch.Tensor((0.25,)), momentum=0.5):
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
seed = extra_args.get("seed", None)
match noise_sampler:
case "brownian":
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
case "gaussian":
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
case "uniform":
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
case "highres-pyramid":
noise_sampler = lambda sigma, sigma_next: highres_pyramid_noise_like(x)
case "perlin":
noise_sampler = lambda sigma, sigma_next: rand_perlin_like(x)
case _:
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
return sample_refined_exp_s(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, noise_sampler=noise_sampler, denoise_to_zero=denoise_to_zero, simple_phi_calc=simple_phi_calc, c2=c2, ita=ita, momentum=momentum)
@torch.no_grad()
def sample_dpmpp_dualsde_momentum(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1/2, momentum=0.0):
"""DPM-Solver++ (Stochastic with Momentum). Personal modified sampler by Clybius"""
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
noise_sampler = rand_perlin_like(x) if noise_sampler is None else noise_sampler
extra_args = {} if extra_args is None else extra_args
s_in = x.new_ones([x.shape[0]])
sigma_fn = lambda t: t.neg().exp()
t_fn = lambda sigma: sigma.log().neg()
denoisedsde_1, denoisedsde_2, denoisedsde_3 = None, None, None # new line
h_1, h_2, h_3 = None, None, None # new line
#def momentum_func(diff, velocity, timescale=1.0, offset=momentum_offset): # Diff is current diff, vel is previous diff
# if velocity is None:
# momentum_vel = diff
# #print("Setting up momentum")
# else:
# momentum_vel = momentum * (timescale - momentum_offset) * velocity + (1 - momentum * (timescale - momentum_offset)) * diff
# #print("Calculating momentum at", momentum)
# return momentum_vel
def momentum_func(diff, velocity, timescale=1.0, offset=-momentum / 2.0): # Diff is current diff, vel is previous diff
if velocity is None:
momentum_vel = diff
else:
momentum_vel = momentum * (timescale + offset) * velocity + (1 - momentum * (timescale + offset)) * diff
return momentum_vel
vel = None
vel_2 = None
vel_sde = None
for i in trange(len(sigmas) - 1, disable=disable):
time = sigmas[i] / sigma_max
denoised = model(x, sigmas[i] * s_in, **extra_args)
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
if sigmas[i + 1] == 0:
# Euler method
d = to_d(x, sigmas[i], denoised)
dt = sigmas[i + 1] - sigmas[i]
x = x + d * dt
else:
# DPM-Solver++
t, t_next = t_fn(sigmas[i]), t_fn(sigmas[i + 1])
h = t_next - t
h_eta = h * (eta + 1)
s = t + h * r
fac = 1 / (2 * r)
# Step 1
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta)
s_ = t_fn(sd)
diff_2 = momentum_func((t - s_).expm1() * denoised, vel_2, time)
vel_2 = diff_2
x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - diff_2
x_2 = x_2 + noise_sampler(sigma_fn(t), sigma_fn(s)) * s_noise * su
denoised_2 = model(x_2, sigma_fn(s) * s_in, **extra_args)
# Step 2
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta)
t_next_ = t_fn(sd)
denoised_d = (1 - fac) * denoised + fac * denoised_2
diff = momentum_func((t - t_next_).expm1() * denoised_d, vel, time)
vel = diff
x = (sigma_fn(t_next_) / sigma_fn(t)) * x - diff
if h_3 is not None:
r0 = h_3 / h_2
r1 = h_2 / h
r2 = h / h_1
d1_0 = (denoised_d - denoisedsde_1) / r2
d1_1 = (denoisedsde_1 - denoisedsde_2) / r1
d1_2 = (denoisedsde_2 - denoisedsde_3) / r0
d1 = d1_0 + (d1_0 - d1_1) * r2 / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) * r2 / ((r2 + r1) * (r0 + r1))
d2 = (d1_0 - d1_1) / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) / ((r2 + r1) * (r0 + r1))
phi_3 = h_eta.neg().expm1() / h_eta + 1
phi_4 = phi_3 / h_eta - 0.5
diff = momentum_func(phi_3 * d1 - phi_4 * d2, vel_sde, time)
vel_sde = diff
x = x + diff
elif h_2 is not None:
r0 = h_1 / h
r1 = h_2 / h
d1_0 = (denoised_d - denoisedsde_1) / r0
d1_1 = (denoisedsde_1 - denoisedsde_2) / r1
d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1)
d2 = (d1_0 - d1_1) / (r0 + r1)
phi_2 = h_eta.neg().expm1() / h_eta + 1
phi_3 = phi_2 / h_eta - 0.5
diff = momentum_func(phi_2 * d1 - phi_3 * d2, vel_sde, time)
vel_sde = diff
x = x + diff
elif h_1 is not None:
r = h_1 / h
d = (denoised_d - denoisedsde_1) / r
phi_2 = h_eta.neg().expm1() / h_eta + 1
diff = momentum_func(phi_2 * d, vel_sde, time)
vel_sde = diff
x = x + diff
if eta:
x = x + noise_sampler(sigma_fn(t), sigma_fn(t_next)) * s_noise * su
#if 'denoised_d' in locals():
denoisedsde_1, denoisedsde_2, denoisedsde_3 = denoised_d, denoisedsde_1, denoisedsde_2 # new line
#if 'h' in locals():
h_1, h_2, h_3 = h, h_1, h_2
return x
def sample_dpmpp_dualsdemomentum(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1/2, momentum=0.0):
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
seed = extra_args.get("seed", None)
match noise_sampler:
case "brownian":
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
case "gaussian":
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
case "uniform":
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
case "perlin":
noise_sampler = lambda sigma, sigma_next: rand_perlin_like(x)
case _:
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
return sample_dpmpp_dualsde_momentum(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, r=r, momentum=momentum)
from .other_samplers.sample_ttm import sample_ttm_jvp
def sample_ttmcustom(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
seed = extra_args.get("seed", None)
match noise_sampler:
case "gaussian":
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
case "uniform":
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
case "brownian":
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
case _:
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
return sample_ttm_jvp(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler)
from k_diffusion.sampling import sample_lcm
def sample_lcmcustom(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None):
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
seed = extra_args.get("seed", None)
match noise_sampler:
case "gaussian":
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
case "uniform":
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
case "brownian":
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
case _:
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
return sample_lcm(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, noise_sampler=noise_sampler)
def sample_clyb_4m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler="brownian", momentum=0.5):
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
seed = extra_args.get("seed", None)
match noise_sampler:
case "brownian":
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
case "gaussian":
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
case "uniform":
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
case "highres-pyramid":
noise_sampler = lambda sigma, sigma_next: highres_pyramid_noise_like(x)
case "perlin":
noise_sampler = lambda sigma, sigma_next: rand_perlin_like(x)
case _:
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
return sample_clyb_4m_sde_momentumized(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, momentum=momentum)
# Add your personal samplers below here, just for formatting purposes ;3
# Add any extra samplers to the following dictionary
extra_samplers = {
"res_momentumized": sample_res_solver,
"dpmpp_dualsde_momentumized": sample_dpmpp_dualsdemomentum,
"clyb_4m_sde_momentumized": sample_clyb_4m_sde,
"ttm": sample_ttmcustom,
"lcm_custom_noise": sample_lcmcustom,
}
+474
View File
@@ -0,0 +1,474 @@
from .other_samplers.refined_exp_solver import sample_refined_exp_s
import comfy.samplers
import comfy.sample
from comfy.k_diffusion import sampling as k_diffusion_sampling
import latent_preview
import torch
import numpy as np
from tqdm.auto import trange
import random
def pyramid_noise_like(size, dtype, layout, generator, device="cpu", discount=0.8):
b, c, h, w = size
orig_h = h
orig_w = w
noise = torch.zeros(size=size, dtype=dtype, layout=layout, device=device)
r = 1
for i in range(5):
r *= 2 # Rather than always going 2x,
#w, h = max(1, int(w/(r**i))), max(1, int(h/(r**i)))
noise += torch.nn.functional.interpolate((torch.normal(mean=0, std=0.5 ** i, size=(b, c, h * r, w * r), dtype=dtype, layout=layout, generator=generator, device=device)), size=(orig_h, orig_w), mode='nearest-exact') * discount**i
#if w>=orig_w*16 or h>=orig_h*16: break
return noise
def power_noise_sampler(size, dtype, layout, generator, device="cpu", alpha=2, k=1): # This doesn't work properly right now
"""Generate 1/f noise for a given tensor.
Args:
tensor: The tensor to add noise to.
alpha: The parameter that determines the slope of the spectrum.
k: A constant.
Returns:
A tensor with the same shape as `tensor` containing 1/f noise.
"""
tensor = torch.randn(size=size, dtype=dtype, layout=layout, generator=generator, device=device)
fft = torch.fft.fft2(tensor)
freq = torch.arange(1, len(fft) + 1, dtype=torch.float)
spectral_density = k / freq**alpha
noise = torch.rand(size=size, dtype=dtype, layout=layout, generator=generator, device=device) * spectral_density
mean = torch.mean(noise, dim=(-2, -1), keepdim=True).to(tensor.device)
std = torch.std(noise, dim=(-2, -1), keepdim=True).to(tensor.device)
noise = noise.to(tensor.device).sub_(mean).div_(std)
return noise
def prepare_noise(latent_image, seed, noise_type, noise_inds=None): # From `sample.py`
"""
creates random noise given a latent image and a seed.
optional arg skip can be used to skip and discard x number of noise generations for a given seed
"""
generator = torch.manual_seed(seed)
match noise_type:
case "gaussian":
noise_func = torch.randn
case "uniform":
def uniform_rand(*size, **kwargs):
return (torch.rand(*size, **kwargs) - 0.5) * 2 * 1.73
noise_func = uniform_rand
case "pyramid":
noise_func = pyramid_noise_like
case "power":
noise_func = power_noise_sampler
case _:
noise_func = torch.randn
if noise_inds is None:
return noise_func(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, generator=generator, device="cpu")
unique_inds, inverse = np.unique(noise_inds, return_inverse=True)
noises = []
for i in range(unique_inds[-1]+1):
noise = noise_func([1] + list(latent_image.size())[1:], dtype=latent_image.dtype, layout=latent_image.layout, generator=generator, device="cpu")
if i in unique_inds:
noises.append(noise)
noises = [noises[i] for i in inverse]
noises = torch.cat(noises, axis=0)
return noises
class SamplerRES_MOMENTUMIZED:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"noise_sampler_type": (["gaussian", "uniform", "brownian", "highres-pyramid", "perlin"], ),
"momentum": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step":0.01}),
"denoise_to_zero": ("BOOLEAN", {"default": True}),
"simple_phi_calc": ("BOOLEAN", {"default": False}),
"ita": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 100.0, "step":0.01, "round": False}),
"c2": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step":0.01, "round": False}),
}
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling"
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type, momentum, denoise_to_zero, simple_phi_calc, ita, c2):
sampler = comfy.samplers.ksampler("res_momentumized", {"noise_sampler": noise_sampler_type, "denoise_to_zero": denoise_to_zero, "simple_phi_calc": simple_phi_calc, "c2": c2, "ita": torch.Tensor((ita,)), "momentum": momentum})
return (sampler, )
class SamplerDPMPP_DUALSDE_MOMENTUMIZED:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"noise_sampler_type": (["gaussian", "uniform", "brownian", "perlin"], ),
"momentum": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step":0.01}),
"eta": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
"s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
"r": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 100.0, "step":0.01}),
}
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling"
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type, momentum, eta, s_noise, r,):
sampler = comfy.samplers.ksampler("dpmpp_dualsde_momentumized", {"noise_sampler": noise_sampler_type, "eta": eta, "s_noise": s_noise, "r": r, "momentum": momentum})
return (sampler, )
class SamplerTTM:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"noise_sampler_type": (["gaussian", "uniform", "brownian"], ),
"eta": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
"s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
}
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling"
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type, eta, s_noise):
sampler = comfy.samplers.ksampler("ttm", {"noise_sampler": noise_sampler_type, "eta": eta, "s_noise": s_noise})
return (sampler, )
class SamplerLCMCustom:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"noise_sampler_type": (["gaussian", "uniform", "brownian"], ),
}
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling"
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type):
sampler = comfy.samplers.ksampler("lcm_custom_noise", {"noise_sampler": noise_sampler_type})
return (sampler, )
class SamplerCLYB_4M_SDE_MOMENTUMIZED:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"noise_sampler_type": (["gaussian", "uniform", "brownian", "highres-pyramid", "perlin"], ),
"momentum": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step":0.01}),
"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
"s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
}
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling"
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type, eta, s_noise, momentum):
sampler = comfy.samplers.ksampler("clyb_4m_sde_momentumized", {"noise_sampler": noise_sampler_type, "eta": eta, "s_noise": s_noise, "momentum": momentum})
return (sampler, )
from comfy import model_management
import comfy.utils
import comfy.conds
from comfy.sample import prepare_sampling, cleanup_additional_models, get_models_from_cond
def mixture_sample(model, model2, noise, positive, positive2, negative, negative2, cfg, cfg2, device, device2, sampler, sampler2, sigmas, sigmas2, model_options={}, model_options2={}, latent_image=None, denoise_mask=None, denoise_mask2=None, callback=None, callback2=None, disable_pbar=False, seed=None):
positive = positive[:]
negative = negative[:]
positive2 = positive2[:]
negative2 = negative2[:]
comfy.samplers.resolve_areas_and_cond_masks(positive, noise.shape[2], noise.shape[3], device)
comfy.samplers.resolve_areas_and_cond_masks(negative, noise.shape[2], noise.shape[3], device)
comfy.samplers.resolve_areas_and_cond_masks(positive2, noise.shape[2], noise.shape[3], device2)
comfy.samplers.resolve_areas_and_cond_masks(negative2, noise.shape[2], noise.shape[3], device2)
model_wrap = comfy.samplers.wrap_model(model)
model_wrap2 = comfy.samplers.wrap_model(model2)
comfy.samplers.calculate_start_end_timesteps(model, negative)
comfy.samplers.calculate_start_end_timesteps(model, positive)
comfy.samplers.calculate_start_end_timesteps(model2, negative2)
comfy.samplers.calculate_start_end_timesteps(model2, positive2)
if latent_image is not None:
latent_image = model.process_latent_in(latent_image)
if hasattr(model, 'extra_conds'):
positive = comfy.samplers.encode_model_conds(model.extra_conds, positive, noise, device, "positive", latent_image=latent_image, denoise_mask=denoise_mask, seed=seed)
negative = comfy.samplers.encode_model_conds(model.extra_conds, negative, noise, device, "negative", latent_image=latent_image, denoise_mask=denoise_mask, seed=seed)
if hasattr(model2, 'extra_conds'):
positive = comfy.samplers.encode_model_conds(model2.extra_conds, positive2, noise, device2, "positive", latent_image=latent_image, denoise_mask=denoise_mask2, seed=seed)
negative = comfy.samplers.encode_model_conds(model2.extra_conds, negative2, noise, device2, "negative", latent_image=latent_image, denoise_mask=denoise_mask2, seed=seed)
#make sure each cond area has an opposite one with the same area
for c in positive:
comfy.samplers.create_cond_with_same_area_if_none(negative, c)
for c in negative:
comfy.samplers.create_cond_with_same_area_if_none(positive, c)
for c in positive2:
comfy.samplers.create_cond_with_same_area_if_none(negative2, c)
for c in negative2:
comfy.samplers.create_cond_with_same_area_if_none(positive2, c)
comfy.samplers.pre_run_control(model, negative + positive)
comfy.samplers.pre_run_control(model2, negative2 + positive2)
comfy.samplers.apply_empty_x_to_equal_area(list(filter(lambda c: c.get('control_apply_to_uncond', False) == True, positive)), negative, 'control', lambda cond_cnets, x: cond_cnets[x])
comfy.samplers.apply_empty_x_to_equal_area(positive, negative, 'gligen', lambda cond_cnets, x: cond_cnets[x])
comfy.samplers.apply_empty_x_to_equal_area(list(filter(lambda c: c.get('control_apply_to_uncond', False) == True, positive2)), negative2, 'control', lambda cond_cnets, x: cond_cnets[x])
comfy.samplers.apply_empty_x_to_equal_area(positive2, negative2, 'gligen', lambda cond_cnets, x: cond_cnets[x])
extra_args = {"cond":positive, "uncond":negative, "cond_scale": cfg, "model_options": model_options, "seed":seed}
extra_args2 = {"cond":positive2, "uncond":negative2, "cond_scale": cfg, "model_options": model_options2, "seed":seed}
samples = None
temp_sigmas = sigmas
temp_sigmas2 = sigmas2
#samples = sampler.sample(model_wrap, sigmas, extra_args, callback, noise, latent_image, denoise_mask, True)
for i in trange(len(sigmas) - 1, disable=disable_pbar):
last_step = i + 1
start_step = i
if last_step is not None and last_step < (len(sigmas) - 1):
temp_sigmas = sigmas[:last_step + 1]
temp_sigmas2 = sigmas2[:last_step + 1]
if start_step is not None:
if start_step < (len(sigmas) - 1):
temp_sigmas = temp_sigmas[start_step:]
temp_sigmas2 = temp_sigmas2[start_step:]
else:
if latent_image is not None:
return latent_image
else:
return torch.zeros_like(noise)
if len(temp_sigmas) != 2:
temp_sigmas = sigmas[-2:]
temp_sigmas2 = sigmas2[-2:]
if (i % 2) == 0:
#print(temp_sigmas)
samples = sampler.sample(model_wrap, temp_sigmas, extra_args, callback, noise.to(device) if i is 0 else torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device=device), samples if samples is not None else latent_image, denoise_mask, True)
else:
#print(temp_sigmas)
samples = sampler2.sample(model_wrap2, temp_sigmas2, extra_args2, callback2, noise.to(device2) if i is 0 else torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device=device2), samples if samples is not None else latent_image, denoise_mask2, True)
return model.process_latent_out(samples.to(torch.float32))
def sample_mixture(model, model2, noise, cfg, cfg2, sampler, sampler2, sigmas, sigmas2, positive, negative, latent_image, noise_mask=None, callback=None, callback2=None, disable_pbar=False, seed=None):
real_model, positive_copy, negative_copy, noise_mask, models = prepare_sampling(model, noise.shape, positive, negative, noise_mask)
real_model2, positive_copy2, negative_copy2, noise_mask2, models2 = prepare_sampling(model2, noise.shape, positive, negative, noise_mask)
noise = noise.to(model.load_device)
latent_image = latent_image.to(model.load_device)
sigmas = sigmas.to(model.load_device)
sigmas2 = sigmas2.to(model.load_device)
samples = mixture_sample(real_model, real_model2, noise, positive_copy, positive_copy2, negative_copy, negative_copy2, cfg, cfg2, model.load_device, model2.load_device, sampler, sampler2, sigmas, sigmas2, model_options=model.model_options, model_options2=model2.model_options, latent_image=latent_image, denoise_mask=noise_mask, denoise_mask2=noise_mask2, callback=callback, callback2=callback2, disable_pbar=disable_pbar, seed=seed)
samples = samples.to(comfy.model_management.intermediate_device())
cleanup_additional_models(models)
cleanup_additional_models(models2)
cleanup_additional_models(set(get_models_from_cond(positive_copy, "control") + get_models_from_cond(negative_copy, "control")))
cleanup_additional_models(set(get_models_from_cond(positive_copy2, "control") + get_models_from_cond(negative_copy2, "control")))
return samples
class SamplerCustomNoise:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"model": ("MODEL",),
"add_noise": ("BOOLEAN", {"default": True}),
"noise_is_latent": ("BOOLEAN", {"default": False}),
"noise_type": (["gaussian", "uniform", "pyramid", "power"], ),
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.5, "round": 0.01}),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"sampler": ("SAMPLER", ),
"sigmas": ("SIGMAS", ),
"latent_image": ("LATENT", ),
}
}
RETURN_TYPES = ("LATENT","LATENT")
RETURN_NAMES = ("output", "denoised_output")
FUNCTION = "sample"
CATEGORY = "sampling/custom_sampling"
def sample(self, model, add_noise, noise_is_latent, noise_type, noise_seed, cfg, positive, negative, sampler, sigmas, latent_image):
latent = latent_image
latent_image = latent["samples"]
if not add_noise:
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
else:
batch_inds = latent["batch_index"] if "batch_index" in latent else None
noise = prepare_noise(latent_image, noise_seed, noise_type, batch_inds)
if noise_is_latent:
noise += latent_image.cpu()# * noise.std()
noise.sub_(noise.mean()).div_(noise.std())
noise_mask = None
if "noise_mask" in latent:
noise_mask = latent["noise_mask"]
x0_output = {}
callback = latent_preview.prepare_callback(model, sigmas.shape[-1] - 1, x0_output)
disable_pbar = False
samples = comfy.sample.sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
out = latent.copy()
out["samples"] = samples
if "x0" in x0_output:
out_denoised = latent.copy()
out_denoised["samples"] = model.model.process_latent_out(x0_output["x0"].cpu())
else:
out_denoised = out
return (out, out_denoised)
class SamplerCustomNoiseDuo:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"model": ("MODEL",),
"add_noise": ("BOOLEAN", {"default": True}),
"add_noise_pass2": ("BOOLEAN", {"default": True}),
"return_noisy_pass1": ("BOOLEAN", {"default": False}),
"noise_type": (["gaussian", "uniform", "pyramid", "power"], ),
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"cfg2": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"sampler": ("SAMPLER", ),
"sampler2": ("SAMPLER", ),
"sigmas": ("SIGMAS", ),
"sigmas2": ("SIGMAS", ),
"hr_upscale": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 9.0, "step":0.1, "round": 0.01}),
"latent_image": ("LATENT", ),
}
}
RETURN_TYPES = ("LATENT","LATENT")
RETURN_NAMES = ("output", "denoised_output")
FUNCTION = "sample"
CATEGORY = "sampling/custom_sampling"
def sample(self, model, add_noise, add_noise_pass2, return_noisy_pass1, noise_type, noise_seed, cfg, cfg2, positive, negative, sampler, sampler2, sigmas, sigmas2, hr_upscale, latent_image):
latent = latent_image
latent_image = latent["samples"]
if not add_noise:
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
else:
batch_inds = latent["batch_index"] if "batch_index" in latent else None
noise = prepare_noise(latent_image, noise_seed, noise_type, batch_inds)
noise_mask = None
if "noise_mask" in latent:
noise_mask = latent["noise_mask"]
x0_output = {}
callback = latent_preview.prepare_callback(model, sigmas.shape[-1] - 1, x0_output)
disable_pbar = False
samples = comfy.sample.sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
if not return_noisy_pass1:
out_denoised = latent.copy()
samples = model.model.process_latent_out(x0_output["x0"].cpu())
if hr_upscale > 1.0:
if "noise_mask" in latent:
noise_mask = comfy.utils.common_upscale(noise_mask, (int)(noise_mask.shape[-1] * hr_upscale), (int)(noise_mask.shape[-2] * hr_upscale), "bislerp", "disabled")
samples = comfy.utils.common_upscale(samples, (int)(samples.shape[-1] * hr_upscale), (int)(samples.shape[-2] * hr_upscale), "bislerp", "disabled")
noise = prepare_noise(samples, noise_seed, noise_type, batch_inds)
samples = comfy.sample.sample_custom(model, noise if add_noise_pass2 else torch.zeros(samples.size(), dtype=samples.dtype, layout=samples.layout, device="cpu"), cfg2, sampler2, sigmas2, positive, negative, samples, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
out = latent.copy()
out["samples"] = samples
if "x0" in x0_output:
out_denoised = latent.copy()
out_denoised["samples"] = model.model.process_latent_out(x0_output["x0"].cpu())
else:
out_denoised = out
return (out, out_denoised)
class SamplerCustomModelMixtureDuo:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"model": ("MODEL",),
"model2": ("MODEL",),
"add_noise": ("BOOLEAN", {"default": True}),
"add_noise_pass2": ("BOOLEAN", {"default": True}),
"return_noisy_pass1": ("BOOLEAN", {"default": False}),
"noise_type": (["gaussian", "uniform", "pyramid", "power"], ),
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"cfg2": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"sampler": ("SAMPLER", ),
"sampler2": ("SAMPLER", ),
"sigmas": ("SIGMAS", ),
"sigmas2": ("SIGMAS", ),
"hr_upscale": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 9.0, "step":0.1, "round": 0.01}),
"latent_image": ("LATENT", ),
}
}
RETURN_TYPES = ("LATENT","LATENT")
RETURN_NAMES = ("output", "denoised_output")
FUNCTION = "sample"
CATEGORY = "sampling/custom_sampling"
def sample(self, model, model2, add_noise, add_noise_pass2, return_noisy_pass1, noise_type, noise_seed, cfg, cfg2, positive, negative, sampler, sampler2, sigmas, sigmas2, hr_upscale, latent_image):
latent = latent_image
latent_image = latent["samples"]
if not add_noise:
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
else:
batch_inds = latent["batch_index"] if "batch_index" in latent else None
noise = prepare_noise(latent_image, noise_seed, noise_type, batch_inds)
noise_mask = None
if "noise_mask" in latent:
noise_mask = latent["noise_mask"]
x0_output = {}
callback = latent_preview.prepare_callback(model, sigmas.shape[-1] - 1, x0_output)
callback2 = latent_preview.prepare_callback(model2, sigmas.shape[-1] - 1, x0_output)
disable_pbar = False
samples = sample_mixture(model, model2, noise, cfg, cfg2, sampler, sampler2, sigmas, sigmas2, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, callback2=callback2, disable_pbar=disable_pbar, seed=noise_seed)
#if not return_noisy_pass1:
# out_denoised = latent.copy()
# samples = model.model.process_latent_out(x0_output["x0"].cpu())
#if hr_upscale > 1.0:
# if "noise_mask" in latent:
# noise_mask = comfy.utils.common_upscale(noise_mask, (int)(noise_mask.shape[-1] * hr_upscale), (int)(noise_mask.shape[-2] * hr_upscale), "bislerp", "disabled")
# samples = comfy.utils.common_upscale(samples, (int)(samples.shape[-1] * hr_upscale), (int)(samples.shape[-2] * hr_upscale), "bislerp", "disabled")
# noise = prepare_noise(samples, noise_seed, noise_type, batch_inds)
#samples = sample_mixture(model, model2, noise if add_noise_pass2 else torch.zeros(samples.size(), dtype=samples.dtype, layout=samples.layout, device="cpu"), cfg2, sampler2, sigmas2, positive, negative, samples, noise_mask=noise_mask, callback=callback, callback2=callback2, disable_pbar=disable_pbar, seed=noise_seed)
out = latent.copy()
out["samples"] = samples
if "x0" in x0_output:
out_denoised = latent.copy()
out_denoised["samples"] = model.model.process_latent_out(x0_output["x0"].cpu())
else:
out_denoised = out
return (out, out_denoised)
+296
View File
@@ -0,0 +1,296 @@
import torch
from torch import no_grad, FloatTensor
from tqdm import tqdm
from itertools import pairwise
from typing import Protocol, Optional, Dict, Any, TypedDict, NamedTuple, Union, List
import math
class DenoiserModel(Protocol):
def __call__(self, x: FloatTensor, t: FloatTensor, *args, **kwargs) -> FloatTensor: ...
class RefinedExpCallbackPayload(TypedDict):
x: FloatTensor
i: int
sigma: FloatTensor
sigma_hat: FloatTensor
class RefinedExpCallback(Protocol):
def __call__(self, payload: RefinedExpCallbackPayload) -> None: ...
class NoiseSampler(Protocol):
def __call__(self, x: FloatTensor) -> FloatTensor: ...
class StepOutput(NamedTuple):
x_next: FloatTensor
denoised: FloatTensor
denoised2: FloatTensor
vel: FloatTensor
vel_2: FloatTensor
def _gamma(
n: int,
) -> int:
"""
https://en.wikipedia.org/wiki/Gamma_function
for every positive integer n,
Γ(n) = (n-1)!
"""
return math.factorial(n-1)
def _incomplete_gamma(
s: int,
x: float,
gamma_s: Optional[int] = None
) -> float:
"""
https://en.wikipedia.org/wiki/Incomplete_gamma_function#Special_values
if s is a positive integer,
Γ(s, x) = (s-1)!*∑{k=0..s-1}(x^k/k!)
"""
if gamma_s is None:
gamma_s = _gamma(s)
sum_: float = 0
# {k=0..s-1} inclusive
for k in range(s):
numerator: float = x**k
denom: int = math.factorial(k)
quotient: float = numerator/denom
sum_ += quotient
incomplete_gamma_: float = sum_ * math.exp(-x) * gamma_s
return incomplete_gamma_
# by Katherine Crowson
def _phi_1(neg_h: FloatTensor):
return torch.nan_to_num(torch.expm1(neg_h) / neg_h, nan=1.0)
# by Katherine Crowson
def _phi_2(neg_h: FloatTensor):
return torch.nan_to_num((torch.expm1(neg_h) - neg_h) / neg_h**2, nan=0.5)
# by Katherine Crowson
def _phi_3(neg_h: FloatTensor):
return torch.nan_to_num((torch.expm1(neg_h) - neg_h - neg_h**2 / 2) / neg_h**3, nan=1 / 6)
def _phi(
neg_h: float,
j: int,
):
"""
For j={1,2,3}: you could alternatively use Kat's phi_1, phi_2, phi_3 which perform fewer steps
Lemma 1
https://arxiv.org/abs/2308.02157
ϕj(-h) = 1/h^j*∫{0..h}(e^(τ-h)*(τ^(j-1))/((j-1)!)dτ)
https://www.wolframalpha.com/input?i=integrate+e%5E%28%CF%84-h%29*%28%CF%84%5E%28j-1%29%2F%28j-1%29%21%29d%CF%84
= 1/h^j*[(e^(-h)*(-τ)^(-j)*τ(j))/((j-1)!)]{0..h}
https://www.wolframalpha.com/input?i=integrate+e%5E%28%CF%84-h%29*%28%CF%84%5E%28j-1%29%2F%28j-1%29%21%29d%CF%84+between+0+and+h
= 1/h^j*((e^(-h)*(-h)^(-j)*h^j*(Γ(j)-Γ(j,-h)))/(j-1)!)
= (e^(-h)*(-h)^(-j)*h^j*(Γ(j)-Γ(j,-h))/((j-1)!*h^j)
= (e^(-h)*(-h)^(-j)*(Γ(j)-Γ(j,-h))/(j-1)!
= (e^(-h)*(-h)^(-j)*(Γ(j)-Γ(j,-h))/Γ(j)
= (e^(-h)*(-h)^(-j)*(1-Γ(j,-h)/Γ(j))
requires j>0
"""
assert j > 0
gamma_: float = _gamma(j)
incomp_gamma_: float = _incomplete_gamma(j, neg_h, gamma_s=gamma_)
phi_: float = math.exp(neg_h) * neg_h**-j * (1-incomp_gamma_/gamma_)
return phi_
class RESDECoeffsSecondOrder(NamedTuple):
a2_1: float
b1: float
b2: float
def _de_second_order(
h: float,
c2: float,
simple_phi_calc = False,
) -> RESDECoeffsSecondOrder:
"""
Table 3
https://arxiv.org/abs/2308.02157
ϕi,j := ϕi,j(-h) = ϕi(-cj*h)
a2_1 = c2ϕ1,2
= c2ϕ1(-c2*h)
b1 = ϕ1 - ϕ2/c2
"""
if simple_phi_calc:
# Kat computed simpler expressions for phi for cases j={1,2,3}
a2_1: float = c2 * _phi_1(-c2*h)
phi1: float = _phi_1(-h)
phi2: float = _phi_2(-h)
else:
# I computed general solution instead.
# they're close, but there are slight differences. not sure which would be more prone to numerical error.
a2_1: float = c2 * _phi(j=1, neg_h=-c2*h)
phi1: float = _phi(j=1, neg_h=-h)
phi2: float = _phi(j=2, neg_h=-h)
phi2_c2: float = phi2/c2
b1: float = phi1 - phi2_c2
b2: float = phi2_c2
return RESDECoeffsSecondOrder(
a2_1=a2_1,
b1=b1,
b2=b2,
)
def _refined_exp_sosu_step(
model: DenoiserModel,
x: FloatTensor,
sigma: FloatTensor,
sigma_next: FloatTensor,
c2 = 0.5,
extra_args: Dict[str, Any] = {},
pbar: Optional[tqdm] = None,
simple_phi_calc = False,
momentum = 0.0,
vel = None,
vel_2 = None,
time = None
) -> StepOutput:
"""
Algorithm 1 "RES Second order Single Update Step with c2"
https://arxiv.org/abs/2308.02157
Parameters:
model (`DenoiserModel`): a k-diffusion wrapped denoiser model (e.g. a subclass of DiscreteEpsDDPMDenoiser)
x (`FloatTensor`): noised latents (or RGB I suppose), e.g. torch.randn((B, C, H, W)) * sigma[0]
sigma (`FloatTensor`): timestep to denoise
sigma_next (`FloatTensor`): timestep+1 to denoise
c2 (`float`, *optional*, defaults to .5): partial step size for solving ODE. .5 = midpoint method
extra_args (`Dict[str, Any]`, *optional*, defaults to `{}`): kwargs to pass to `model#__call__()`
pbar (`tqdm`, *optional*, defaults to `None`): progress bar to update after each model call
simple_phi_calc (`bool`, *optional*, defaults to `True`): True = calculate phi_i,j(-h) via simplified formulae specific to j={1,2}. False = Use general solution that works for any j. Mathematically equivalent, but could be numeric differences.
"""
def momentum_func(diff, velocity, timescale=1.0, offset=-momentum / 2.0): # Diff is current diff, vel is previous diff
if velocity is None:
momentum_vel = diff
else:
momentum_vel = momentum * (timescale + offset) * velocity + (1 - momentum * (timescale + offset)) * diff
return momentum_vel
lam_next, lam = (s.log().neg() for s in (sigma_next, sigma))
# type hints aren't strictly true regarding float vs FloatTensor.
# everything gets promoted to `FloatTensor` after interacting with `sigma: FloatTensor`.
# I will use float to indicate any variables which are scalars.
h: float = lam_next - lam
a2_1, b1, b2 = _de_second_order(h=h, c2=c2, simple_phi_calc=simple_phi_calc)
denoised: FloatTensor = model(x, sigma, **extra_args)
if pbar is not None:
pbar.update(0.5)
c2_h: float = c2*h
diff_2 = momentum_func(a2_1*h*denoised, vel_2, time)
vel_2 = diff_2
x_2: FloatTensor = math.exp(-c2_h)*x + diff_2
lam_2: float = lam + c2_h
sigma_2: float = lam_2.neg().exp()
denoised2: FloatTensor = model(x_2, sigma_2, **extra_args)
if pbar is not None:
pbar.update(0.5)
diff = momentum_func(h*(b1*denoised + b2*denoised2), vel, time)
vel = diff
x_next: FloatTensor = math.exp(-h)*x + diff
return StepOutput(
x_next=x_next,
denoised=denoised,
denoised2=denoised2,
vel=vel,
vel_2=vel_2,
)
@no_grad()
def sample_refined_exp_s(
model: FloatTensor,
x: FloatTensor,
sigmas: FloatTensor,
denoise_to_zero: bool = True,
extra_args: Dict[str, Any] = {},
callback: Optional[RefinedExpCallback] = None,
disable: Optional[bool] = None,
ita: FloatTensor = torch.zeros((1,)),
c2 = .5,
noise_sampler: NoiseSampler = torch.randn_like,
simple_phi_calc = False,
momentum = 0.0,
):
"""
Refined Exponential Solver (S).
Algorithm 2 "RES Single-Step Sampler" with Algorithm 1 second-order step
https://arxiv.org/abs/2308.02157
Parameters:
model (`DenoiserModel`): a k-diffusion wrapped denoiser model (e.g. a subclass of DiscreteEpsDDPMDenoiser)
x (`FloatTensor`): noised latents (or RGB I suppose), e.g. torch.randn((B, C, H, W)) * sigma[0]
sigmas (`FloatTensor`): sigmas (ideally an exponential schedule!) e.g. get_sigmas_exponential(n=25, sigma_min=model.sigma_min, sigma_max=model.sigma_max)
denoise_to_zero (`bool`, *optional*, defaults to `True`): whether to finish with a first-order step down to 0 (rather than stopping at sigma_min). True = fully denoise image. False = match Algorithm 2 in paper
extra_args (`Dict[str, Any]`, *optional*, defaults to `{}`): kwargs to pass to `model#__call__()`
callback (`RefinedExpCallback`, *optional*, defaults to `None`): you can supply this callback to see the intermediate denoising results, e.g. to preview each step of the denoising process
disable (`bool`, *optional*, defaults to `False`): whether to hide `tqdm`'s progress bar animation from being printed
ita (`FloatTensor`, *optional*, defaults to 0.): degree of stochasticity, η, for each timestep. tensor shape must be broadcastable to 1-dimensional tensor with length `len(sigmas) if denoise_to_zero else len(sigmas)-1`. each element should be from 0 to 1.
c2 (`float`, *optional*, defaults to .5): partial step size for solving ODE. .5 = midpoint method
noise_sampler (`NoiseSampler`, *optional*, defaults to `torch.randn_like`): method used for adding noise
simple_phi_calc (`bool`, *optional*, defaults to `True`): True = calculate phi_i,j(-h) via simplified formulae specific to j={1,2}. False = Use general solution that works for any j. Mathematically equivalent, but could be numeric differences.
"""
#assert sigmas[-1] == 0
ita = ita.to(x.device)
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
vel, vel_2 = None, None
with tqdm(disable=disable, total=len(sigmas)-(1 if denoise_to_zero else 2)) as pbar:
for i, (sigma, sigma_next) in enumerate(pairwise(sigmas[:-1].split(1))):
time = sigmas[i] / sigma_max
if 'sigma' not in locals():
sigma = sigmas[i]
eps = noise_sampler(sigma, sigma_next).float()
sigma_hat = sigma * (1 + ita)
x_hat = x + (sigma_hat ** 2 - sigma ** 2) ** .5 * eps
x_next, denoised, denoised2, vel, vel_2 = _refined_exp_sosu_step(
model,
x_hat,
sigma_hat,
sigma_next,
c2=c2,
extra_args=extra_args,
pbar=pbar,
simple_phi_calc=simple_phi_calc,
momentum = momentum,
vel = vel,
vel_2 = vel_2,
time = time
)
if callback is not None:
payload = RefinedExpCallbackPayload(
x=x,
i=i,
sigma=sigma,
sigma_hat=sigma_hat,
denoised=denoised,
denoised2=denoised2,
)
callback(payload)
x = x_next
if denoise_to_zero:
eps = noise_sampler(sigma, sigma_next).float()
sigma_hat = sigma * (1 + ita)
x_hat = x + (sigma_hat ** 2 - sigma ** 2) ** .5 * eps
x_next: FloatTensor = model(x_hat, sigma.to(x_hat.device), **extra_args)
pbar.update()
x = x_next
return x
+118
View File
@@ -0,0 +1,118 @@
from k_diffusion.sampling import default_noise_sampler, to_d
from tqdm import trange
import torch
from torch import enable_grad
# by Katherine Crowson
@enable_grad()
def sample_ttm(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
import torch.autograd.forward_ad as fwAD
extra_args = {} if extra_args is None else extra_args
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
s_in = x.new_ones([x.shape[0]])
for i in trange(len(sigmas) - 1, disable=disable):
denoised = model(x, sigmas[i] * s_in, **extra_args)
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
if sigmas[i + 1] == 0:
# Denoising step
x = denoised
else:
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
h = s - t
h_eta = h * (eta + 1)
with fwAD.dual_level():
eps = to_d(x, sigmas[i], denoised)
dual_x = fwAD.make_dual(x, eps * -sigmas[i])
dual_sigma = fwAD.make_dual(sigmas[i] * s_in, -sigmas[i] * s_in)
dual_denoised = model(dual_x, dual_sigma, **extra_args)
denoised_prime = fwAD.unpack_dual(dual_denoised).tangent
phi_1 = -torch.expm1(-h_eta)
phi_2 = torch.expm1(-h_eta) + h_eta
x = torch.exp(-h_eta) * x + phi_1 * denoised + phi_2 * denoised_prime
if eta:
phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * eta))
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * phi_1_noise * s_noise
return x
# by Katherine Crowson
def sample_ttm_jvp(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
"""Second order truncated Taylor method (torch.func.jvp() version)."""
extra_args = {} if extra_args is None else extra_args
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
s_in = x.new_ones([x.shape[0]])
model_fn = lambda x, sigma: model(x, sigma * s_in, **extra_args)
for i in trange(len(sigmas) - 1, disable=disable):
denoised = model_fn(x, sigmas[i])
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
if sigmas[i + 1] == 0:
# Denoising step
x = denoised
else:
# 2nd order truncated Taylor method
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
h = s - t
h_eta = h * (eta + 1)
eps = to_d(x, sigmas[i], denoised)
_, denoised_prime = torch.func.jvp(model_fn, (x, sigmas[i]), (eps * -sigmas[i], -sigmas[i]))
phi_1 = -torch.expm1(-h_eta)
phi_2 = torch.expm1(-h_eta) + h_eta
# phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0
x = torch.exp(-h_eta) * x + phi_1 * denoised + phi_2 * denoised_prime
if eta:
phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * eta))
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * phi_1_noise * s_noise
return x
@torch.no_grad()
def sample_lcm_ttm_jvp(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None, eta=0.0):
extra_args = {} if extra_args is None else extra_args
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
s_in = x.new_ones([x.shape[0]])
model_fn = lambda x, sigma: model(x, sigma * s_in, **extra_args)
for i in trange(len(sigmas) - 1, disable=disable):
denoised = model_fn(x, sigmas[i])
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
if sigmas[i + 1] == 0:
# Denoising step
x = denoised
else:
# 2nd order truncated Taylor method
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
h = s - t
h_eta = h * (eta + 1)
x = denoised
x_2 = x + sigmas[i + 1] * noise_sampler(sigmas[i], sigmas[i + 1])
eps = to_d(x_2, sigmas[i + 1], denoised)
_, denoised_prime = torch.func.jvp(model_fn, (x_2, sigmas[i + 1]), (eps * -sigmas[i + 1], -sigmas[i + 1]))
phi_1 = -torch.expm1(-h_eta)
phi_2 = torch.expm1(-h_eta) + h_eta
# phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0
x = torch.exp(-h_eta) * x + phi_1 * denoised + phi_2 * denoised_prime
if sigmas[i + 1] > 0:
x = x + sigmas[i + 1] * noise_sampler(sigmas[i], sigmas[i + 1])
return x