diff --git a/restart_sampling.py b/restart_sampling.py index d2b972a..7aff720 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -4,10 +4,12 @@ import ast import os import warnings from collections import namedtuple +from typing import NamedTuple import comfy import latent_preview import torch +from comfy import model_sampling from comfy.sample import prepare_noise, sample_custom from comfy.samplers import KSAMPLER, KSampler, sampler_object from comfy.utils import ProgressBar @@ -570,6 +572,49 @@ class RestartPlan: print("\n|| Done test") +class RestartScaleFactors(NamedTuple): + latent_scale: float = 1.0 + noise_scale: float = 0.0 + + @classmethod + def build( + cls, + sigma_from: torch.Tensor, + sigma_to: torch.Tensor, + *, + is_flow: bool, + ): + sigma_from = sigma_from.detach().cpu().item() + sigma_to = sigma_to.detach().cpu().item() + if sigma_from == sigma_to: + return cls(latent_scale=1.0, noise_scale=0.0) + if sigma_from > sigma_to: + raise ValueError("Can't do a restart step to a lower sigma!") + + if not is_flow: + return cls( + latent_scale=1.0, + noise_scale=max(0.0, (sigma_to**2 - sigma_from**2)) ** 0.5, + ) + + snr_to = 1.0 - sigma_to + if snr_to <= 0: + # It may make more sense to clamp the noise scale to 1.0. + return cls(latent_scale=0.0, noise_scale=max(1.0, sigma_to)) + snr_from = 1.0 - sigma_from + latent_scale = snr_to / snr_from + noise_scale = max(0.0, (sigma_to**2) - (latent_scale * sigma_from) ** 2) ** 0.5 + return cls(latent_scale, noise_scale) + + def add_noise(self, latent: torch.Tensor, noise: torch.Tensor) -> torch.Tensor: + if self.latent_scale != 1.0: + latent *= self.latent_scale + if self.noise_scale != 0: + noise *= self.noise_scale + latent += noise + return latent + + class RestartSampler: @staticmethod def get_segment(sigmas: torch.Tensor) -> torch.Tensor: @@ -584,22 +629,22 @@ class RestartSampler: return sigmas @classmethod - def split_sigmas(cls, sigmas): + def split_sigmas(cls, sigmas, *, is_flow: bool=False): # 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). + # It's a generator that yields tuples of (RestartScaleFactors, 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]: s_min, s_max = prev_seg[-1], seg[0] - noise_scale = ((s_max**2 - s_min**2) ** 0.5).item() + scale_factors = RestartScaleFactors.build(s_min, s_max, is_flow=is_flow) else: - noise_scale = 0.0 + scale_factors = RestartScaleFactors() prev_seg = seg - yield (noise_scale, seg) + yield (scale_factors, seg) # Some extra explanation for a couple of these arguments: # @@ -638,9 +683,11 @@ class RestartSampler: if restart_custom_noise is not None: restart_noise = restart_custom_noise + is_flow = isinstance(model.inner_model.inner_model.model_sampling, model_sampling.CONST) + sampler = restart_wrapped_sampler.sampler_function - chunks = tuple(cls.split_sigmas(sigmas)) + chunks = tuple(cls.split_sigmas(sigmas, is_flow=is_flow)) total_steps = sum(len(chunk) - 1 for _noise, chunk in chunks) step = 0 noise_count = 0 @@ -676,12 +723,20 @@ class RestartSampler: **kwargs, ) - for noise_scale, chunk_sigmas in chunks: - if noise_scale != 0: + for scale_factors, chunk_sigmas in chunks: + if scale_factors.noise_scale != 0: s_min, s_max = chunk_sigmas[-1], chunk_sigmas[0] - x += ( - restart_noise(x, s_min, s_max, seed + noise_count)(s_max, s_min) - * noise_scale + x = scale_factors.add_noise( + x, + restart_noise( + x, + s_min, + s_max, + seed + noise_count, + )( + s_max, + s_min, + ), ) noise_count += 1 if restart_chunked: