* Weird experiments with sampler blending * Mega sync * Use ComfyUI union types for wildcard inputs when available * Add weoon wavelet sampler * Fix noise sampler caching not considering immiscible settings Allow specifying alt custom noise/immiscible settings for samplers that internally add noise AFS support at the merge sampler level Allow the Weoon internal step to use ETA Add t_copysign expression function * Start updating documentation * Add ImmiscibleReference node, other changes * Respect factor in ImmiscibleReference noise * Fixing normalizing in ImmiscibleReference noise * Handle case where there are no RGB factors (i.e. audio models) * Fix rectified flow type detection for non-Flux models * Add more tensor ops, add OCS ApplyExpressionImage/Latent nodes * Rename ApplyExpression nodes to ApplyFilter, expression QoL improvements * Fix apply filter nodes definitions * Make t_noise handler work for image-latents. * Add t_new_like expression function * Make it possible to comment out lines with hash in expressions * Add gradient estimation, pingpong and res_multistep samplers. Add pingpong group merge method. Remove model call caching stuff. Allow defining pre_cfg and post_cfg filters. Handle filters changing the latent shape better. Allow disabling detecting out of order sigmas as restart sampling. * Sync current updates (which I am too lazy to describe individually) * Add ExpressionFilteredNoise node Improvements to immiscible noise handling/expanded features
357 lines
11 KiB
Python
357 lines
11 KiB
Python
import torch
|
|
|
|
import comfy
|
|
from comfy.k_diffusion.sampling import to_d
|
|
|
|
from . import filtering
|
|
|
|
from .utils import fallback
|
|
from .latent import OCSLatentFormat
|
|
|
|
|
|
class History:
|
|
def __init__(self, size):
|
|
self.history = []
|
|
self.size = size
|
|
|
|
def __len__(self):
|
|
return len(self.history)
|
|
|
|
def __getitem__(self, k):
|
|
return self.history[k]
|
|
|
|
def push(self, val):
|
|
if len(self.history) >= self.size:
|
|
self.history = self.history[-(self.size - 1) :]
|
|
self.history.append(val)
|
|
|
|
def reset(self):
|
|
self.history = []
|
|
|
|
def clone(self):
|
|
obj = self.__new__(self.__class__)
|
|
obj.__init__(self.size)
|
|
obj.history = self.history.copy()
|
|
return obj
|
|
|
|
|
|
class ModelResult:
|
|
def __init__(
|
|
self,
|
|
call_idx,
|
|
sigma,
|
|
x,
|
|
denoised,
|
|
have_uncond=True,
|
|
**kwargs,
|
|
):
|
|
self.call_idx = call_idx
|
|
self.sigma = sigma
|
|
self.x = x
|
|
self.denoised = denoised
|
|
self.have_uncond = have_uncond
|
|
for k in ("denoised_uncond", "denoised_cond", "tangents", "jdenoised"):
|
|
setattr(self, k, kwargs.pop(k, None))
|
|
if len(kwargs) != 0:
|
|
raise ValueError(f"Unexpected keyword arguments: {tuple(kwargs.keys())}")
|
|
|
|
def to_d(
|
|
self,
|
|
/,
|
|
x=None,
|
|
sigma=None,
|
|
denoised=None,
|
|
denoised_uncond=None,
|
|
alt_cfgpp_scale=0,
|
|
cfgpp=False,
|
|
):
|
|
x = fallback(x, self.x)
|
|
sigma = fallback(sigma, self.sigma)
|
|
denoised = fallback(denoised, self.denoised)
|
|
denoised_uncond = fallback(denoised_uncond, self.denoised_uncond)
|
|
if alt_cfgpp_scale != 0:
|
|
x = x - denoised * alt_cfgpp_scale + denoised_uncond * alt_cfgpp_scale
|
|
return to_d(x, sigma, denoised if not cfgpp else denoised_uncond)
|
|
|
|
def get_split_prediction(
|
|
self,
|
|
*,
|
|
x=None,
|
|
d=None,
|
|
sigma=None,
|
|
denoised=None,
|
|
denoised_uncond=None,
|
|
alt_cfgpp_scale=0,
|
|
cfgpp=False,
|
|
):
|
|
denoised = fallback(denoised, self.denoised)
|
|
denoised_uncond = fallback(denoised_uncond, self.denoised_uncond)
|
|
x = fallback(x, self.x)
|
|
sigma = fallback(sigma, self.sigma)
|
|
if d is None:
|
|
d = self.to_d(
|
|
x=x,
|
|
sigma=sigma,
|
|
denoised=denoised,
|
|
denoised_uncond=denoised_uncond,
|
|
alt_cfgpp_scale=alt_cfgpp_scale,
|
|
cfgpp=cfgpp,
|
|
)
|
|
denoised_pred = denoised if alt_cfgpp_scale == 0 else x - d * sigma
|
|
return (denoised_pred, d)
|
|
|
|
@property
|
|
def d(self):
|
|
return self.to_d()
|
|
|
|
def clone(self, deep=False):
|
|
obj = self.__new__(self.__class__)
|
|
for k in (
|
|
"denoised",
|
|
"call_idx",
|
|
"sigma",
|
|
"x",
|
|
"denoised_uncond",
|
|
"denoised_cond",
|
|
"tangents",
|
|
"jdenoised",
|
|
):
|
|
val = getattr(self, k)
|
|
if deep and isinstance(val, torch.Tensor):
|
|
val = val.copy()
|
|
setattr(obj, k, val)
|
|
return obj
|
|
|
|
def get_error(self, other, *, override=None, alt_cfgpp_scale=0, cfgpp=False):
|
|
slf = fallback(override, self)
|
|
first, second = (other, slf) if other.sigma > slf.sigma else (slf, other)
|
|
if first.sigma == second.sigma:
|
|
return 0.0
|
|
d = first.to_d(alt_cfgpp_scale=alt_cfgpp_scale, cfgpp=cfgpp)
|
|
d_pred = second.to_d(
|
|
x=first.x + d * (second.sigma - first.sigma),
|
|
alt_cfgpp_scale=alt_cfgpp_scale,
|
|
cfgpp=cfgpp,
|
|
)
|
|
return torch.linalg.norm(d_pred.sub_(d)).div_(torch.linalg.norm(d)).item()
|
|
|
|
|
|
class OCSModel:
|
|
def __init__(
|
|
self,
|
|
model,
|
|
x: torch.Tensor,
|
|
s_in: torch.Tensor,
|
|
extra_args: dict,
|
|
*,
|
|
cache: dict | None = None,
|
|
filter: dict | None = None,
|
|
cfg1_uncond_optimization: bool = False,
|
|
cfg_scale_override: int | float | None = None,
|
|
) -> None:
|
|
filtargs = fallback(filter, {}).copy()
|
|
self.filters = {}
|
|
for key in (
|
|
"input",
|
|
"denoised",
|
|
"jdenoised",
|
|
"cond",
|
|
"uncond",
|
|
"postcfg",
|
|
"precfg",
|
|
):
|
|
filt = filtargs.pop(key, None)
|
|
if filt is None:
|
|
continue
|
|
self.filters[key] = filtering.make_filter(filt)
|
|
self.model = model
|
|
self.s_in = s_in
|
|
self.extra_args = extra_args
|
|
self.cfg1_uncond_optimization = cfg1_uncond_optimization
|
|
self.cfg_scale_override = cfg_scale_override
|
|
self.is_rectified_flow = isinstance(
|
|
model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST
|
|
)
|
|
self.latent_format = OCSLatentFormat(
|
|
x.device, model.inner_model.inner_model.latent_format
|
|
)
|
|
|
|
def maybe_filter(
|
|
self, name: str, latent: torch.Tensor, *args: list, **kwargs: dict
|
|
) -> torch.Tensor:
|
|
filt = self.filters.get(name)
|
|
if filt is None:
|
|
return latent
|
|
return filt.apply(latent, *args, **kwargs)
|
|
|
|
def filter_result(
|
|
self, result: ModelResult, *args: list, **kwargs: dict
|
|
) -> ModelResult:
|
|
if not self.filters:
|
|
return result
|
|
result = result.clone()
|
|
for key in ("denoised", "cond", "uncond", "jdenoised"):
|
|
filt = self.filters.get(key)
|
|
if filt is None:
|
|
continue
|
|
attk = f"denoised_{key}" if key in {"cond", "uncond"} else key
|
|
inpval = getattr(result, attk, None)
|
|
if inpval is None:
|
|
continue
|
|
setattr(result, attk, filt.apply(inpval, *args, **kwargs))
|
|
return result
|
|
|
|
@staticmethod
|
|
def _fr_add_mr(fr: filtering.FilterRefs, mr: ModelResult) -> filtering.FilterRefs:
|
|
frmr = filtering.FilterRefs.from_mr(mr)
|
|
fr.kvs |= {f"{k}_curr": v for k, v in frmr.kvs.items()}
|
|
return fr
|
|
|
|
def call_model(
|
|
self, x: torch.Tensor, sigma: torch.Tensor, **kwargs: dict
|
|
) -> torch.Tensor:
|
|
return self.model(x, sigma * self.s_in, **self.extra_args | kwargs)
|
|
|
|
@property
|
|
def model_sampling(self):
|
|
return self.model.inner_model.inner_model.model_sampling
|
|
|
|
@property
|
|
def inner_cfg_scale(self) -> None | int | float:
|
|
maybe_cfg_scale = getattr(self.model.inner_model, "cfg", None)
|
|
return maybe_cfg_scale if isinstance(maybe_cfg_scale, (int, float)) else None
|
|
|
|
def set_inner_cfg_scale(self, scale: None | int | float) -> None | int | float:
|
|
eff_scale = self.cfg_scale_override
|
|
if scale is not None:
|
|
eff_scale = None if scale < 0 else scale
|
|
if eff_scale is None or eff_scale < 0:
|
|
return None
|
|
curr_cfg_scale = self.inner_cfg_scale
|
|
if curr_cfg_scale is None:
|
|
return None
|
|
self.model.inner_model.cfg = eff_scale
|
|
return curr_cfg_scale
|
|
|
|
def __call__(
|
|
self,
|
|
x: torch.Tensor,
|
|
sigma: torch.Tensor,
|
|
*,
|
|
call_index: int = 0,
|
|
ss,
|
|
s_in=None,
|
|
tangents=None,
|
|
require_uncond: bool = False,
|
|
cfg_scale_override: None | int = None,
|
|
**kwargs,
|
|
) -> ModelResult:
|
|
filter_refs = ss.refs | filtering.FilterRefs({
|
|
"model_call": call_index,
|
|
"orig_x": x,
|
|
})
|
|
|
|
comfy.model_management.throw_exception_if_processing_interrupted()
|
|
|
|
model_options = self.extra_args.get("model_options", {}).copy()
|
|
denoised_cond = denoised_uncond = None
|
|
have_uncond = True
|
|
|
|
def postcfg(args):
|
|
nonlocal denoised_cond, denoised_uncond, have_uncond
|
|
denoised_uncond = args["uncond_denoised"]
|
|
denoised_cond = args["cond_denoised"]
|
|
result = args["denoised"]
|
|
if "postcfg" in self.filters:
|
|
result = self.maybe_filter(
|
|
"postcfg",
|
|
result,
|
|
refs=filter_refs
|
|
| filtering.FilterRefs({
|
|
"postcfg_input": args["input"],
|
|
"postcfg_sigma": args["sigma"],
|
|
"postcfg_denoised_cond": denoised_cond,
|
|
"postcfg_denoised_uncond": denoised_uncond,
|
|
}),
|
|
)
|
|
if denoised_uncond is None:
|
|
have_uncond = False
|
|
denoised_uncond = denoised_cond
|
|
return result
|
|
|
|
def precfg(args):
|
|
conds_out = args["conds_out"]
|
|
precfg_refs = (
|
|
filter_refs
|
|
| filtering.FilterRefs({
|
|
"precfg_input": args["input"],
|
|
"precfg_sigma": args["sigma"],
|
|
"precfg_cond_scale": args["cond_scale"],
|
|
})
|
|
| filtering.FilterRefs({
|
|
f"precfg_cond_{idx}": cond for idx, cond in enumerate(conds_out)
|
|
})
|
|
)
|
|
return [
|
|
self.maybe_filter(
|
|
"precfg",
|
|
curr_cond,
|
|
refs=precfg_refs | filtering.FilterRefs({"cond_idx": cond_idx}),
|
|
)
|
|
for cond_idx, curr_cond in enumerate(conds_out)
|
|
]
|
|
|
|
orig_cfg_scale = self.set_inner_cfg_scale(cfg_scale_override)
|
|
|
|
model_options = comfy.model_patcher.set_model_options_post_cfg_function(
|
|
model_options,
|
|
postcfg,
|
|
disable_cfg1_optimization=require_uncond
|
|
or not self.cfg1_uncond_optimization,
|
|
)
|
|
if "precfg" in self.filters:
|
|
model_options = comfy.model_patcher.set_model_options_pre_cfg_function(
|
|
model_options,
|
|
precfg,
|
|
)
|
|
|
|
extra_args = self.extra_args | {"model_options": model_options}
|
|
s_in = fallback(s_in, self.s_in)
|
|
if s_in.shape[0] != x.shape[0]:
|
|
s_in = self.s_in = x.new_ones((x.shape[0],))
|
|
x = self.maybe_filter("input", x, refs=filter_refs)
|
|
|
|
def call_model(x, sigma, **kwargs):
|
|
return self.model(x, sigma * s_in, **extra_args | kwargs)
|
|
|
|
if tangents is None:
|
|
denoised = call_model(x, sigma, **kwargs)
|
|
self.set_inner_cfg_scale(orig_cfg_scale)
|
|
mr = ModelResult(
|
|
call_index,
|
|
sigma,
|
|
x,
|
|
denoised,
|
|
have_uncond=have_uncond,
|
|
denoised_uncond=denoised_uncond,
|
|
denoised_cond=denoised_cond,
|
|
)
|
|
self._fr_add_mr(filter_refs, mr)
|
|
mr = self.filter_result(mr, default_ref=x, refs=filter_refs)
|
|
return mr
|
|
denoised, denoised_prime = torch.func.jvp(call_model, (x, sigma), tangents)
|
|
self.set_inner_cfg_scale(orig_cfg_scale)
|
|
mr = ModelResult(
|
|
call_index,
|
|
sigma,
|
|
x,
|
|
denoised,
|
|
have_uncond=have_uncond,
|
|
jdenoised=denoised_prime,
|
|
denoised_uncond=denoised_uncond,
|
|
denoised_cond=denoised_cond,
|
|
)
|
|
self._fr_add_mr(filter_refs, mr)
|
|
mr = self.filter_result(mr, default_ref=x, refs=filter_refs)
|
|
return mr
|