Merge pull request #26 from blepping/feat_flow_restarting
Support flow models
This commit is contained in:
+66
-11
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user