From 471b5972c99e00817d6531c937b2af019f0d959e Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 18 Apr 2024 10:56:55 -0600 Subject: [PATCH] only generate restart segment down to model sigma_min merge restart segments more aggressively --- restart_sampling.py | 155 +++++++++++++++++++++++++------------------- 1 file changed, 88 insertions(+), 67 deletions(-) diff --git a/restart_sampling.py b/restart_sampling.py index 9306583..53bd463 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -194,7 +194,9 @@ def restart_sampling( ### UNCOMMENT TO RUN SELF TEST # plan.self_test( # model, + # min_steps=2, # schedules=SCHEDULER_MAPPING.keys(), + # # schedules=("simple_test",), # restart_schedules=SCHEDULER_MAPPING.keys(), # ) sigmas = plan.sigmas() @@ -291,45 +293,64 @@ def restart_sampling( class PlanItem( namedtuple( "PlanItem", - ["sigmas", "k", "s_min", "s_max", "restart_sigmas"], - defaults=[None, 0, 0.0, 0.0, None], + ["sigmas", "k", "restart_sigmas"], + defaults=[None, 0, None], ), ): __slots__ = () def __new__(cls, *args: list, **kwargs: dict): - threshold = 1e-06 obj = super().__new__(cls, *args, **kwargs) - if len(obj.sigmas) < 2: + obj.validate() + # print(">>", obj) + return obj + + def validate(self, threshold=1e-06): + if len(self.sigmas) < 2: raise ValueError("PlanItem: invalid normal sigmas: too short") - if obj.k < 1: - return obj - if len(obj.restart_sigmas) < 2: + if self.k < 1: + return + if len(self.restart_sigmas) < 2: raise ValueError("PlanItem: invalid restart sigmas: too short") - if obj.s_min >= obj.s_max: + if self.s_min >= self.s_max: raise ValueError("PlanItem: invalid min/max: min >= max") - if obj.sigmas[-1] - obj.restart_sigmas[0] > threshold: + if self.sigmas[-1] - self.restart_sigmas[0] > threshold: raise ValueError( "PlanItem: invalid sigmas: last normal sigma >= first restart sigma", ) - if obj.sigmas[-1] - obj.restart_sigmas[-1] > threshold: + if self.sigmas[-1] - self.restart_sigmas[-1] > threshold: errstr = ( - f"PlanItem: invalid sigmas: last restart sigma {obj.restart_sigmas[-1]} < last normal sigma {obj.sigmas[-1]}", + f"PlanItem: invalid sigmas: last restart sigma {self.restart_sigmas[-1]} < last normal sigma {self.sigmas[-1]}", ) raise ValueError(errstr) - t = obj.sigmas.sort(descending=True, stable=True)[0].unique_consecutive() - if not torch.equal(obj.sigmas, t): - raise ValueError( - "PlanItem: invalid normal sigmas: out of order or contains duplicates", + t = self.sigmas.sort(descending=True, stable=True)[0].unique_consecutive() + if not torch.equal(self.sigmas, t): + errstr = ( + f"PlanItem: invalid normal sigmas: out of order or contains duplicates: {self}", ) - t = obj.restart_sigmas.sort(descending=True, stable=True)[ + raise ValueError(errstr) + t = self.restart_sigmas.sort(descending=True, stable=True)[ 0 ].unique_consecutive() - if not torch.equal(obj.restart_sigmas, t): - raise ValueError( - "PlanItem: invalid restart sigmas: out of order or contains duplicates", + if not torch.equal(self.restart_sigmas, t): + errstr = ( + f"PlanItem: invalid normal sigmas: out of order or contains duplicates: {self}", ) - return obj + raise ValueError(errstr) + + @property + def total_steps(self): + if self.k < 1: + return len(self.sigmas) - 1 + return (len(self.sigmas) - 1) + (len(self.restart_sigmas) - 1) * self.k + + @property + def s_min(self): + return None if self.k < 1 else self.restart_sigmas[-1].item() + + @property + def s_max(self): + return None if self.k < 1 else self.restart_sigmas[0].item() # Execute a plan item: runs sampling on the main sigmas, handles injecting noise for restarts # as well as sampling the restart steps. @@ -339,8 +360,9 @@ class PlanItem( # It takes x, and sigma_min, sigma_max (basically the same arguments as ComfyUI's # BrownianTreeNoiseSampler class init function). @torch.no_grad() - def execute(self, x, sample, get_noise_sampler): - x = sample(x, self.sigmas, -1) + def execute(self, x, sample, get_noise_sampler, skip_normal=False, next_pi=None): + if not skip_normal: + x = sample(x, self.sigmas, -1) if self.k < 1 or self.restart_sigmas is None: return x noise_sampler = get_noise_sampler(x, self.s_min, self.s_max) @@ -349,7 +371,14 @@ class PlanItem( noise_sampler(self.restart_sigmas[0], self.restart_sigmas[-1]) * (self.s_max**2 - self.s_min**2) ** 0.5 ) - x = sample(x, self.restart_sigmas, kidx) + if next_pi and kidx == self.k - 1: + sigmas = torch.cat((self.restart_sigmas[:-1], next_pi.sigmas)).to( + self.restart_sigmas.device, + ) + print("COMBINE", sigmas) + else: + sigmas = self.restart_sigmas + x = sample(x, sigmas, kidx) return x @@ -366,10 +395,11 @@ class RestartPlan: force_full_denoise=False, sigmas=None, ): - comfy.model_management.load_models_gpu([model]) - real_model = model - while hasattr(real_model, "model"): - real_model = real_model.model + # comfy.model_management.load_models_gpu([model]) + # real_model = model + # while hasattr(real_model, "model"): + # real_model = real_model.model + ms = model.get_model_object("model_sampling") effective_steps = ( steps @@ -380,13 +410,13 @@ class RestartPlan: sigmas = calc_sigmas( scheduler, effective_steps, - float(real_model.model_sampling.sigma_min), - float(real_model.model_sampling.sigma_max), - real_model, + float(ms.sigma_min), + float(ms.sigma_max), + model.model, "cpu", ) else: - sigmas = sigmas.detach().cpu().clone() + sigmas = sigmas.clone().detach().cpu() if step_range is not None: start_step, last_step = step_range @@ -402,11 +432,7 @@ class RestartPlan: self.plain_sigmas = sigmas - restart_segments = prepare_restart_segments( - restart_info, - real_model.model_sampling, - sigmas, - ) + restart_segments = prepare_restart_segments(restart_info, ms, sigmas) self.plan, self.total_steps = self.build_plan_items( model.model, restart_segments, @@ -434,7 +460,6 @@ class RestartPlan: device, ) -> tuple[list, int]: segments = round_restart_segments(sigmas, restart_segments) - total_steps = len(sigmas) - 1 + calc_restart_steps(segments) plan = [] range_start = -1 for i in range(len(sigmas) - 1): @@ -449,22 +474,27 @@ class RestartPlan: if k < 1 or n_restart < 2: continue normal_sigmas = sigmas[range_start : i + 2] + effsmin = max(float(model.model_sampling.sigma_min), sigmas[i + 1]) + if effsmin >= s_max: + continue restart_sigmas = calc_sigmas( restart_scheduler, n_restart, - s_min if restart_scheduler != "exponential" else max(s_min, 1e-08), + effsmin, s_max, model, device=device, - )[:-1] - # restart_sigmas[0] = s_max + ) + if normal_sigmas[-1] != 0: + restart_sigmas = restart_sigmas[:-1] restart_sigmas[-1] = s_min - plan.append(PlanItem(normal_sigmas, k, s_min, s_max, restart_sigmas)) + # restart_sigmas[0] = s_max + plan.append(PlanItem(normal_sigmas, k, restart_sigmas)) range_start = -1 if range_start != -1: # Include sigmas after the last restart segments in the plan. plan.append(PlanItem(sigmas[range_start:])) - return plan, total_steps + return plan, sum(pi.total_steps for pi in plan) @classmethod def from_sigmas(cls, sigmas, threshold=1e-06): @@ -495,9 +525,8 @@ class RestartPlan: last_sigma = sigma raise ValueError("Unexpected end of sigmas in a restart segment") - plain_sigmas = sigmas.detach().cpu().clone() + plain_sigmas = sigmas.clone().detach().cpu() plan = [] - total_steps = 0 while len(sigmas) > 0: # Get the normal segment - a restart segment can never be first. normal_sigmas = get_normal_segment(sigmas) @@ -508,7 +537,6 @@ class RestartPlan: "Encountered invalid normal segment rebuilding sigmas: too short", ) sigmas = sigmas[nslen:] - total_steps += nslen - 1 if len(sigmas) == 0: # No restart segments follow the normal segment so we're done. plan.append(PlanItem(normal_sigmas)) @@ -527,19 +555,10 @@ class RestartPlan: while len(sigmas) > 0 and torch.equal(sigmas[:rslen], restart_sigmas): k += 1 sigmas = sigmas[rslen:] - total_steps += (rslen - 1) * k - plan.append( - PlanItem( - normal_sigmas, - k, - normal_sigmas[-1], - restart_sigmas[0], - restart_sigmas, - ), - ) + plan.append(PlanItem(normal_sigmas, k, restart_sigmas)) obj = cls.__new__(cls) obj.plan = plan - obj.total_steps = total_steps + obj.total_steps = sum(pi.total_steps for pi in plan) obj.plain_sigmas = plain_sigmas return obj @@ -562,15 +581,7 @@ class RestartPlan: if pi.k < 1: items.append(PlanItem(sigmas)) continue - items.append( - PlanItem( - sigmas, - pi.k, - pi.s_min, - pi.s_max, - pi.restart_sigmas.to(device), - ), - ) + items.append(PlanItem(sigmas, pi.k, pi.restart_sigmas.to(device))) return obj # Dumps information about the plan to the console. It uses the normal plan execute @@ -721,8 +732,18 @@ class RestartPlan: return x # Execute the plan items in sequence. - for pi in plan: - x = pi.execute(x, do_sample, get_noise_sampler) + skip = False + for idx in range(len(plan)): + pi = plan[idx] + nextpi = plan[idx + 1] if idx < len(plan) - 1 else None + x = pi.execute( + x, + do_sample, + get_noise_sampler, + skip_normal=skip, + next_pi=nextpi, + ) + skip = pi.k > 0 and len(pi.restart_sigmas) > 2 and nextpi is not None return x @staticmethod