Initial implementation of ancestral sampling for rectified flow models, supports most samplers

This commit is contained in:
blepping
2024-09-01 17:24:56 -06:00
parent 1618093e22
commit cc30fde46c
4 changed files with 54 additions and 17 deletions
+3
View File
@@ -124,6 +124,9 @@ class ModelCallCache:
self.s_in = s_in
self.extra_args = extra_args
self.cfg1_uncond_optimization = cfg1_uncond_optimization
self.is_rectified_flow = x.shape[1] == 16 and isinstance(
model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST
)
if self.cache.size < 1:
return
self.reset_cache()
+26 -14
View File
@@ -77,6 +77,7 @@ class SamplerResult:
noise_sampler=None,
final=True,
):
self.is_rectified_flow = ss.model.is_rectified_flow
self.sampler = sampler
self.sigma_up = fallback(sigma_up, ss.sigma.new_zeros(1))
self.s_noise = fallback(s_noise, sampler.s_noise)
@@ -130,8 +131,12 @@ class SamplerResult:
x = fallback(x, self.x)
if self.sigma_next == 0 or self.noise_scale == 0:
return x
x = x + self.get_noise(ss=ss) * scale
return x
noise = self.get_noise(ss=ss) * scale
if not self.is_rectified_flow:
return x + noise
x_coeff = (1 - self.sigma_next) / (1 - self.sigma_down)
# print(f"\nRF noise: {x_coeff}")
return x_coeff * x + noise
def clone(self):
obj = self.__new__(self.__class__)
@@ -230,7 +235,14 @@ class SingleStepSampler(CFGPPStepMixin):
def euler_step(self, x, ss):
sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss))
d = self.to_d(ss.hcur)
return (yield from self.result(ss, ss.denoised + d * sigma_down, sigma_up))
return (
yield from self.result(
ss,
ss.denoised + d * sigma_down,
sigma_up,
sigma_down=sigma_down,
)
)
def denoised_result(self, ss, **kwargs):
return (
@@ -538,7 +550,7 @@ class ReversibleHeunStep(ReversibleSingleStepSampler):
+ (sigma_down * (d + d_next) / 2)
- correction * reversible_scale
)
yield from self.result(ss, x, sigma_up)
yield from self.result(ss, x, sigma_up, sigma_down=sigma_down)
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
@@ -583,7 +595,7 @@ class ReversibleHeun1SStep(ReversibleSingleStepSampler):
# Update the sample using the Reversible Heun formula
correction = dtr**2 * (d_next - d_prev) / 4
x = x + (dt * (d_prev + d_next) / 2) - correction * reversible_scale
yield from self.result(ss, x, su)
yield from self.result(ss, x, su, sigma_down=sd)
# def __step(self, x, ss):
# if ss.sigma_next == 0:
@@ -703,7 +715,7 @@ class RESStep(SingleStepSampler):
denoised2 = ss.model(x_2, sigma_2, ss=ss, call_index=1).denoised
x = math.exp(-h) * eff_x + h * (b1 * denoised + b2 * denoised2)
yield from self.result(ss, x, sigma_up)
yield from self.result(ss, x, sigma_up, sigma_down=sigma_down)
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
@@ -730,7 +742,7 @@ class TrapezoidalStep(SingleStepSampler):
# Update the sample using the Trapezoidal rule
x = x + dt_2 * (d_i + d_next) / 2
yield from self.result(ss, x, sigma_up)
yield from self.result(ss, x, sigma_up, sigma_down=sigma_down)
class TrapezoidalCycleStep(CycleSingleStepSampler):
@@ -795,7 +807,7 @@ class BogackiStep(ReversibleSingleStepSampler):
# Update the sample
x = (x + 2 * k1 / 9 + k2 / 3 + 4 * k3 / 9) - correction * reversible_scale
yield from self.result(ss, x, su)
yield from self.result(ss, x, su, sigma_down=sd)
class ReversibleBogackiStep(BogackiStep):
@@ -824,7 +836,7 @@ class RK4Step(SingleStepSampler):
# Update the sample
x = x + (k1 + 2 * k2 + 2 * k3 + k4) / 6
yield from self.result(ss, x, sigma_up)
yield from self.result(ss, x, sigma_up, sigma_down=sigma_down)
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
@@ -895,7 +907,7 @@ class EulerDancingStep(SingleStepSampler):
d_2 = to_d(x, sigma_leap, ss.denoised)
dt_2 = sigma_down2 - sigma_leap
x = x + d_2 * dt_2
yield from self.result(ss, x, sigma_up2)
yield from self.result(ss, x, sigma_up2, sigma_down=sigma_down2)
def _step(self, x, ss):
eta = self.get_dyn_eta(ss)
@@ -1073,7 +1085,7 @@ class DPMPP2SStep(SingleStepSampler, DPMPPStepMixin):
x_2 = (sigma_fn(s) / sigma_fn(t)) * eff_x - (-h * r).expm1() * ss.denoised
denoised_2 = ss.model(x_2, sigma_fn(s), ss=ss, call_index=1).denoised
x = (sigma_fn(t_next) / sigma_fn(t)) * eff_x - (-h).expm1() * denoised_2
yield from self.result(ss, x, sigma_up)
yield from self.result(ss, x, sigma_up, sigma_down=sigma_down)
class DPMPPSDEStep(SingleStepSampler, DPMPPStepMixin):
@@ -1116,7 +1128,7 @@ class DPMPPSDEStep(SingleStepSampler, DPMPPStepMixin):
x = (sigma_fn(t_next_) / sigma_fn(t)) * eff_x - (
t - t_next_
).expm1() * denoised_d
yield from self.result(ss, x, su)
yield from self.result(ss, x, su, sigma_down=sd)
# Based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
@@ -2084,7 +2096,7 @@ class HeunStep(ReversibleSingleStepSampler):
result = hcur.denoised + d * s
result += (dt * (d + d_next)) * 0.5
result -= self.reversible_correction(ss, d, d_next)
yield from self.result(ss, result, su)
yield from self.result(ss, result, su, sigma_down=sd)
class Heun1SStep(HeunStep):
@@ -2105,7 +2117,7 @@ class Heun1SStep(HeunStep):
result = hcur.denoised + hcur.sigma * self.to_d(hcur)
result += (dt * (d_prev + d)) * 0.5
result -= self.reversible_correction(ss, d_prev, d)
yield from self.result(ss, result, su)
yield from self.result(ss, result, su, sigma_down=sd)
class AdapterStep(SingleStepSampler):
+1 -1
View File
@@ -133,7 +133,7 @@ class SimpleSubstepsSampler(MergeSubstepsSampler):
ss.refs = FilterRefs.from_ss(ss, have_current=True)
self.callback()
sr = self.simple_substep(x, ssampler)
return self.merge_steps(sr.x, noise=sr.get_noise(ss=ss))
return self.merge_steps(sr.noise_x(ss=ss))
class NormalMergeSubstepsSampler(MergeSubstepsSampler):
+24 -2
View File
@@ -175,16 +175,38 @@ class SamplerState:
self.refs = FilterRefs.from_ss(self)
def get_ancestral_step(self, eta=1.0, sigma=None, sigma_next=None):
sigma = self.sigma if sigma is None else sigma
sigma_next = self.sigma_next if sigma_next is None else sigma_next
if self.model.is_rectified_flow:
return self.get_ancestral_step_rf(
eta=eta, sigma=sigma, sigma_next=sigma_next
)
sigma = fallback(sigma, self.sigma)
sigma_next = fallback(sigma_next, self.sigma_next)
if eta == 0 or sigma_next <= 0:
return sigma_next, sigma_next.new_zeros(1)
sd, su = (
v if isinstance(v, torch.Tensor) else sigma.new_full((1,), v)
for v in get_ancestral_step(
sigma, sigma_next, eta=eta if sigma_next != 0 else 0
)
)
if not (sd > 0 and su > 0):
return sigma_next, sigma_next.new_zeros(1)
return sd, su
# Referenced from Comfy dpmpp_2s_ancestral_RF
def get_ancestral_step_rf(self, eta=1.0, sigma=None, sigma_next=None):
sigma = fallback(sigma, self.sigma)
sigma_next = fallback(sigma_next, self.sigma_next)
if eta == 0 or sigma_next <= 0:
return sigma_next, sigma_next.new_zeros(1)
sigma_down = sigma_next * (1 + (sigma_next / sigma - 1) * eta)
alpha_ip1, alpha_down = 1 - sigma_next, 1 - sigma_down
sigma_up = (sigma_next**2 - sigma_down**2 * alpha_ip1**2 / alpha_down**2) ** 0.5
if not (sigma_down > 0 and sigma_up > 0):
return sigma_next, sigma_next.new_zeros(1)
# print(f"\nRF ancestral: down={sigma_down}, up={sigma_up}")
return sigma_down, sigma_up
def clone_edit(self, **kwargs):
obj = self.__class__.__new__(self.__class__)
for k in self.CLONE_KEYS: