Files
blepping-comfyui_overly_com…/py/substep_sampling.py
T

367 lines
11 KiB
Python

from typing import NamedTuple
import torch
from comfy.k_diffusion.sampling import get_ancestral_step
from .filtering import FilterRefs
from .model import History
from .utils import fallback
class AncestralRatios(NamedTuple):
alpha_t: torch.Tensor
alpha_s: torch.Tensor
sigma_up: torch.Tensor
sigma_down: torch.Tensor
class Items:
def __init__(self, items=None):
self.items = [] if items is None else items
def clone(self):
return self.__class__(items=self.items.copy())
def append(self, item):
self.items.append(item)
return item
def __getitem__(self, key):
return self.items[key]
def __setitem__(self, key, value):
self.items[key] = value
def __len__(self):
return len(self.items)
def __iter__(self):
return self.items.__iter__()
class CommonOptionsItems(Items):
def __init__(self, *, s_noise=1.0, eta=1.0, items=None, **kwargs):
super().__init__(items=items)
self.options = kwargs
self.s_noise = s_noise
self.eta = eta
def clone(self):
obj = super().clone()
obj.options = self.options.copy()
obj.s_noise = self.s_noise
obj.eta = self.eta
return obj
class StepSamplerChain(CommonOptionsItems):
def __init__(
self,
*,
merge_method="divide",
time_mode="step",
time_start=0,
time_end=999,
**kwargs,
):
super().__init__(**kwargs)
self.merge_method = merge_method
if time_mode not in ("step", "step_pct", "sigma"):
raise ValueError("Bad time mode")
self.time_mode = time_mode
self.time_start, self.time_end = time_start, time_end
def clone(self):
obj = super().clone()
obj.merge_method = self.merge_method
obj.time_mode = self.time_mode
obj.time_start, obj.time_end = self.time_start, self.time_end
obj.options = self.options.copy()
return obj
class ParamGroup(Items):
pass
class StepSamplerGroups(CommonOptionsItems):
pass
class SamplerState:
CLONE_KEYS = (
"cfg_scale_override",
"model",
"hist",
"extra_args",
"disable_status",
"eta",
"reta",
"s_noise",
"sigmas",
"callback_",
"noise_sampler",
"noise",
"idx",
"total_steps",
"step",
"substep",
"sigma",
"sigma_next",
"sigma_prev",
"sigma_down",
"sigma_up",
"refs",
)
def __init__(
self,
model,
sigmas,
idx,
extra_args,
*,
step=0,
substep=0,
noise_sampler,
callback=None,
denoised=None,
noise=None,
eta=1.0,
reta=1.0,
s_noise=1.0,
disable_status=False,
history_size=4,
cfg_scale_override=None,
):
self.model = model
self.hist = History(max(1, history_size))
self.extra_args = extra_args
self.eta = eta
self.reta = reta
self.s_noise = s_noise
self.sigmas = sigmas
self.callback_ = callback
self.noise_sampler = noise_sampler
self.noise = noise
self.disable_status = disable_status
self.step = 0
self.substep = 0
self.total_steps = len(sigmas) - 1
self.cfg_scale_override = cfg_scale_override
self.is_flow = self.model.is_rectified_flow
self.offset_sigma = (
model.model_sampling.percent_to_sigma(1e-04) if self.is_flow else None
)
self.update(idx) # Sets idx, sigma_prev, sigma, sigma_down, refs
@property
def hcur(self):
return self.hist[-1]
@property
def hprev(self):
return self.hist[-2]
@property
def denoised(self):
return self.hcur.denoised
@property
def denoised_uncond(self):
return self.hcur.denoised_uncond
@property
def denoised_cond(self):
return self.hcur.denoised_cond
@property
def dt(self):
return self.sigma_next - self.sigma
@property
def d(self):
return self.hcur.d
# These two functions referenced from ComfyUI.
def sigma_to_half_log_snr(
self, *, sigma: torch.Tensor | None = None, idx: int | None = None
) -> torch.Tensor:
if sigma is None and idx is None:
sigma = self.sigma
else:
sigma = sigma if sigma is not None else self.sigmas[idx]
if not self.is_flow:
return sigma.log().neg_()
if sigma.max() >= 1.0:
sigma = sigma * 0.0 + self.offset_sigma
return sigma.logit().neg_()
def half_log_snr_to_sigma(self, half_log_snr: torch.Tensor) -> torch.Tensor:
return (torch.sigmoid if self.is_flow else torch.exp)(half_log_snr.neg())
def update(self, idx=None, step=None, substep=None):
idx = self.idx if idx is None else idx
self.idx = idx
self.sigma_prev = None if idx < 1 else self.sigmas[idx - 1]
self.sigma, self.sigma_next = self.sigmas[idx], self.sigmas[idx + 1]
self.sigma_down, self.sigma_up = get_ancestral_step(
self.sigma, self.sigma_next, eta=self.eta
)
if step is not None:
self.step = step
if substep is not None:
self.substep = substep
self.refs = FilterRefs.from_ss(self)
def get_ancestral_step_ext(
self,
*,
sigma: torch.Tensor | None = None,
sigma_next: torch.Tensor | None = None,
eta: float = 1.0,
retry_increment: int = 0,
):
sigma = fallback(sigma, self.sigma)
sigma_next = fallback(sigma_next, self.sigma_next)
sigma_empty = sigma_next * 0.0
def get_noeta_ratios():
return AncestralRatios(
alpha_t=sigma_empty + 1.0,
alpha_s=sigma_empty + 1.0,
sigma_up=sigma_empty.clone(),
sigma_down=sigma_next.clone(),
)
if eta <= 0 or sigma_next.max().item() <= 1e-08:
return get_noeta_ratios()
orig_dtype = sigma.dtype
sigma = sigma.to(dtype=torch.float64)
sigma_next = sigma_next.to(dtype=torch.float64)
alpha_s = sigma * self.sigma_to_half_log_snr(sigma=sigma).exp()
alpha_t = sigma_next * self.sigma_to_half_log_snr(sigma=sigma_next).exp()
adj_sigma = sigma / alpha_s
adj_sigma_next = sigma_next / alpha_t
sd = su = None
while eta > 0:
sd, su = (
v if isinstance(v, torch.Tensor) else sigma.new_full((1,), v)
for v in get_ancestral_step(adj_sigma, adj_sigma_next, eta=eta)
)
if sd > 0 and su > 0:
break
else:
sd = su = None
if retry_increment <= 0:
break
# print(f"\nETA {eta} failed, retrying with {eta - retry_increment}")
eta -= retry_increment
if sd is None or su is None:
return get_noeta_ratios()
sd = alpha_t * sd
return AncestralRatios(
alpha_t=alpha_t.to(dtype=orig_dtype),
alpha_s=alpha_s.to(dtype=orig_dtype),
sigma_up=su.to(dtype=orig_dtype),
sigma_down=sd.to(dtype=orig_dtype),
)
def get_ancestral_step(
self, eta=1.0, sigma=None, sigma_next=None, retry_increment=0
):
if self.model.is_rectified_flow:
return self.get_ancestral_step_rf(
eta=eta,
sigma=sigma,
sigma_next=sigma_next,
retry_increment=retry_increment,
)
sigma = fallback(sigma, self.sigma)
sigma_next = fallback(sigma_next, self.sigma_next)
if eta <= 0 or sigma_next <= 0:
return sigma_next, sigma_next.new_zeros(1)
while eta > 0:
sd, su = (
v if isinstance(v, torch.Tensor) else sigma.new_full((1,), v)
for v in get_ancestral_step(
sigma, sigma_next, eta=eta if sigma_next != 0 else 0
)
)
if sd > 0 and su > 0:
return sd, su
if retry_increment <= 0:
break
# print(f"\nETA {eta} failed, retrying with {eta - retry_increment}")
eta -= retry_increment
return sigma_next, sigma_next.new_zeros(1)
# Referenced from Comfy dpmpp_2s_ancestral_RF
def get_ancestral_step_rf(
self, eta=1.0, sigma=None, sigma_next=None, retry_increment=0
):
sigma = fallback(sigma, self.sigma)
sigma_next = fallback(sigma_next, self.sigma_next)
if eta <= 0 or sigma_next <= 0:
return sigma_next, sigma_next.new_zeros(1)
while eta > 0:
sigma_down = sigma_next * (1 + (sigma_next / sigma - 1) * eta)
alpha_ip1, alpha_down = 1 - sigma_next, 1 - sigma_down
sigma_up = (
sigma_next**2 - sigma_down**2 * alpha_ip1**2 / alpha_down**2
) ** 0.5
if sigma_down > 0 and sigma_up > 0:
return sigma_down, sigma_up
if retry_increment <= 0:
break
eta -= retry_increment
return sigma_next, sigma_next.new_zeros(1)
# print(f"\nRF ancestral: down={sigma_down}, up={sigma_up}")
def clone_edit(self, **kwargs):
obj = self.__class__.__new__(self.__class__)
for k in self.CLONE_KEYS:
setattr(obj, k, kwargs[k] if k in kwargs else getattr(self, k))
obj.update()
return obj
def callback(self, hi=None, *, preview_mode="denoised"):
if not self.callback_:
return None
hi = self.hcur if hi is None else hi
if preview_mode == "cond":
preview = fallback(hi.denoised_cond, hi.denoised)
elif preview_mode == "uncond":
preview = fallback(hi.denoised_uncond, hi.denoised)
elif preview_mode == "raw":
preview = hi.x
elif (
preview_mode == "diff"
and hi.denoised_uncond is not None
and hi.denoised_cond is not None
):
preview = (
hi.denoised_uncond * 0.25 + (hi.denoised_uncond - hi.denoised_cond) * 16
)
elif preview_mode == "noisy":
preview = (hi.x - hi.denoised) * 0.1 + hi.denoised
else:
preview = hi.denoised
return self.callback_(
{
"x": hi.x,
"i": self.step,
"sigma": hi.sigma,
"sigma_hat": hi.sigma,
"denoised": preview,
}
)
def reset(self):
self.hist.reset()
self.denoised = None
def call_model(self, *args, **kwargs):
cfg_scale_override = kwargs.pop("cfg_scale_override", self.cfg_scale_override)
return self.model(*args, cfg_scale_override=cfg_scale_override, **kwargs)