From a36442dd4f3c9ecb835b2df48ef1b41d26f419fa Mon Sep 17 00:00:00 2001 From: blepping Date: Sun, 21 Apr 2024 04:56:59 -0600 Subject: [PATCH] Cleanups, add some comments --- restart_sampling.py | 87 ++++++++++++++++++++------------------------- 1 file changed, 38 insertions(+), 49 deletions(-) diff --git a/restart_sampling.py b/restart_sampling.py index 6d3cb9e..40773d4 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -76,6 +76,7 @@ def prepare_restart_segments(restart_info, ms, sigmas): raise temp = [] default_segments = ast.literal_eval(DEFAULT_SEGMENTS) + # This phase expands any preset strings into actual 4-item restart segments. for idx in range(len(restart_arrays)): item = restart_arrays[idx] if not isinstance(item, str): @@ -90,6 +91,7 @@ def prepare_restart_segments(restart_info, ms, sigmas): raise ValueError("Ill-formed restart segment") restart_arrays = temp restart_segments = [] + # Now we build the actual restart segments. for arr in restart_arrays: if not isinstance(arr, (list, tuple)) or len(arr) != 4: raise ValueError("Restart segment must be a list with 4 values") @@ -148,13 +150,6 @@ def calc_sigmas(scheduler, n, sigma_min, sigma_max, model, device): return SCHEDULER_MAPPING[scheduler](model, n, sigma_min, sigma_max, device) -def calc_restart_steps(restart_segments): - restart_steps = 0 - for segment in restart_segments.values(): - restart_steps += (segment["n"] - 1) * segment["k"] - return restart_steps - - def restart_sampling( model, seed, @@ -294,14 +289,21 @@ class PlanItem( def __new__(cls, *args: list, **kwargs: dict): obj = super().__new__(cls, *args, **kwargs) 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 self.k < 1: + 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}", + ) + raise ValueError(errstr) + if self.k == 0: return + if self.k < 0: + raise ValueError("PlanItem: invalid negative k value") if len(self.restart_sigmas) < 2: raise ValueError("PlanItem: invalid restart sigmas: too short") if self.s_min >= self.s_max: @@ -315,18 +317,12 @@ class PlanItem( f"PlanItem: invalid sigmas: last restart sigma {self.restart_sigmas[-1]} < last normal sigma {self.sigmas[-1]}", ) raise ValueError(errstr) - 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}", - ) - raise ValueError(errstr) t = self.restart_sigmas.sort(descending=True, stable=True)[ 0 ].unique_consecutive() if not torch.equal(self.restart_sigmas, t): errstr = ( - f"PlanItem: invalid normal sigmas: out of order or contains duplicates: {self}", + f"PlanItem: invalid restart sigmas: out of order or contains duplicates: {self}", ) raise ValueError(errstr) @@ -414,6 +410,7 @@ class RestartPlan: device, ) -> tuple[list, int]: model_sigma_min = float(model.model_sampling.sigma_min) + model_sigma_max = float(model.model_sampling.sigma_max) segments = round_restart_segments(sigmas, restart_segments) plan = [] range_start = -1 @@ -436,14 +433,13 @@ class RestartPlan: restart_scheduler, n_restart, max(model_sigma_min, sigmas[i + 1]), - s_max, + min(model_sigma_max, s_max), model, device=device, ) if normal_sigmas[-1] != 0: restart_sigmas = restart_sigmas[:-1] - restart_sigmas[-1] = s_min - # restart_sigmas[0] = s_max + restart_sigmas[-1] = s_min # Force the restart segment to end at s_min. plan.append(PlanItem(normal_sigmas, k, restart_sigmas)) range_start = -1 if range_start != -1: @@ -452,6 +448,8 @@ class RestartPlan: return plan, sum(pi.total_steps for pi in plan) def sigmas(self) -> torch.Tensor: + # Flattens a plan into sigmas. When the first normal sigma matches the last item's + # final sigma, we strip the first normal sigma to avoid creating duplicates. def sigmas_generator(): prev_last = None for pi in self.plan: @@ -547,15 +545,17 @@ class RestartSampler: @classmethod def split_sigmas(cls, sigmas): + # This function just splits the sigmas into chunks that are sorted descending. + # If the first sigma of a chunk is > the last sigma of the previous chunk then this + # is a restart segment: noising the restart uses s_min=prev_chunk[-1], s_max=chunk[0]. + # It's a generator that yields tuples of (noise_scale, chunk_sigmas). prev_seg = None while len(sigmas) > 1: seg = cls.get_segment(sigmas) sigmas = sigmas[len(seg) :] if prev_seg is not None and seg[0] > prev_seg[-1]: - print( - f"CALC NOISE: min={prev_seg[-1].item():.04}, max={seg[0].item():.04}", - ) - noise_scale = ((seg[0] ** 2 - prev_seg[-1] ** 2) ** 0.5).item() + s_min, s_max = prev_seg[-1], seg[0] + noise_scale = ((s_max**2 - s_min**2) ** 0.5).item() else: noise_scale = 0.0 prev_seg = seg @@ -600,7 +600,6 @@ class RestartSampler: sampler = restart_wrapped_sampler.sampler_function - print("SAMPLING", sigmas) total_steps = len(sigmas - 1) step = 0 noise_count = 0 @@ -625,38 +624,28 @@ class RestartSampler: if callback is not None: callback(cb_state) + def do_sample(x, sigmas): + return sampler( + model, + x, + sigmas, + *args, + callback=cb_wrapper, + disable=True, + **kwargs, + ) + for noise_scale, chunk_sigmas in cls.split_sigmas(sigmas): - print(f"CHUNK: noise={noise_scale:.04}, sigmas={chunk_sigmas}") if noise_scale != 0: + s_min, s_max = chunk_sigmas[-1], chunk_sigmas[0] x += ( - restart_noise( - x, - chunk_sigmas[-1], - chunk_sigmas[0], - seed + noise_count, - )(chunk_sigmas[0], chunk_sigmas[-1]) + restart_noise(x, s_min, s_max, seed + noise_count)(s_max, s_min) * noise_scale ) noise_count += 1 if restart_chunked: - x = sampler( - model, - x, - chunk_sigmas, - *args, - callback=cb_wrapper, - disable=True, - **kwargs, - ) + x = do_sample(x, chunk_sigmas) continue for i in range(len(chunk_sigmas) - 1): - x = sampler( - model, - x, - chunk_sigmas[i : i + 2], - *args, - callback=cb_wrapper, - disable=True, - **kwargs, - ) + x = do_sample(x, chunk_sigmas[i : i + 2]) return x