From cc30fde46c2a94ca7aa352590be7f0f34ddff488 Mon Sep 17 00:00:00 2001 From: blepping Date: Sun, 1 Sep 2024 17:24:56 -0600 Subject: [PATCH] Initial implementation of ancestral sampling for rectified flow models, supports most samplers --- py/model.py | 3 +++ py/step_samplers.py | 40 ++++++++++++++++++++++++++-------------- py/substep_merging.py | 2 +- py/substep_sampling.py | 26 ++++++++++++++++++++++++-- 4 files changed, 54 insertions(+), 17 deletions(-) diff --git a/py/model.py b/py/model.py index ca9879c..6bc1111 100644 --- a/py/model.py +++ b/py/model.py @@ -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() diff --git a/py/step_samplers.py b/py/step_samplers.py index f71473b..4d1a934 100644 --- a/py/step_samplers.py +++ b/py/step_samplers.py @@ -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): diff --git a/py/substep_merging.py b/py/substep_merging.py index ab7df52..8b231f2 100644 --- a/py/substep_merging.py +++ b/py/substep_merging.py @@ -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): diff --git a/py/substep_sampling.py b/py/substep_sampling.py index aea8cdf..7c9608f 100644 --- a/py/substep_sampling.py +++ b/py/substep_sampling.py @@ -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: