Merge pull request #26 from blepping/feat_flow_restarting

Support flow models
This commit is contained in:
ssitu
2025-12-05 22:42:31 -05:00
committed by GitHub
+66 -11
View File
@@ -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: