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

358 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.model_sampling = model.inner_model.inner_model.model_sampling
self.is_rectified_flow = isinstance(
self.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