Files
blepping-ComfyUI-bleh/py/nodes/modelPatchConditional.py
T
blepping c98990a3cb Add BlehBlockCFG node
Add tooltips and descriptions for most nodes

Improvements to TAESD previews

Various cleanups and lint squashing

Add more scaling and blend types

Hopefully improve latent normalization (may change seeds)
2024-08-30 11:21:26 -06:00

360 lines
13 KiB
Python

from __future__ import annotations
from typing import Any
import torch
from comfy.ldm.modules.attention import optimized_attention
# ** transformer_options
# input_block_patch* : p(h, transformer_options) -> h
# input_block_patch_after_skip*: p(h, transformer_options) -> h
# output_block_patch* : p(h, hsp, transformer_options) -> h, hsp
# attn1_patch* : p(n, context_attn1, value_attn1, extra_options) -> n, context_attn1, value_attn1
# attn1_output_patch* : p(n, extra_options) -> n
# attn2_patch* : p(n, context_attn2, value_attn2, extra_options) -> n, context_attn2, value_attn2
# attn2_output_patch* : p(n, extra_options) -> n
# middle_patch* : p(x, extra_options) -> x
# attn1_replace* : p(n, context_attn1, value_attn1, extra_options) -> n
# attn2_replace* : p(n, context_attn2, value_attn2, extra_options) -> n
# ** model_options
# model_function_wrapper : p(model.apply_model, {"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond}) -> output
# sampler_cfg_function : p({"cond": x - cond_pred, "uncond": x - uncond_pred, "cond_scale": cond_scale, "timestep": timestep, "input": x, "sigma": timestep, "cond_denoised": cond_pred, "uncond_denoised": uncond_pred, "model": model, "model_options": model_options}) -> output
# sampler_post_cfg_function* : p({"denoised": cfg_result, "cond": cond, "uncond": uncond, "model": model, "uncond_denoised": uncond_pred, "cond_denoised": cond_pred, "sigma": timestep, "model_options": model_options, "input": x}) -> output
class PatchTypeTransformer:
def __init__(
self,
name,
nresult=1,
):
self.name = name
self.nresult = nresult
def get_patches(self, model_options):
return (
model_options.get("transformer_options", {})
.get("patches", {})
.get(self.name, [])
)
def set_patches(self, model_options, val):
to = model_options.get("transformer_options", {})
model_options["transformer_options"] = to
patches = to.get("patches", {})
to["patches"] = patches
patches[self.name] = val
def exists(self, model_options):
return len(self.get_patches(model_options)) > 0
def _call(self, patches, *args: list[Any]):
result_part, arg_part = args[: self.nresult], args[self.nresult :]
for p in patches:
result_part = p(*result_part, *arg_part)
return result_part
@torch.no_grad()
def __call__(self, model_options, *args: list[Any]):
result = self._call(self.get_patches(model_options), *args)
return (
result[0]
if isinstance(result, (tuple, list)) and len(result) == 1
else result
)
class PatchTypeTransformerReplace(PatchTypeTransformer):
def get_patches(self, model_options):
return (
model_options.get("transformer_options", {})
.get("patches_replace", {})
.get(self.name, {})
)
def set_patches(self, model_options, val):
to = model_options.get("transformer_options", {})
model_options["transformer_options"] = to
patches = to.get("patches_replace", {})
to["patches_replace"] = patches
patches[self.name] = val
def __call__(self, key, model_options, *args: list[Any]):
return self._call(key, self.get_patches(model_options), *args)
@classmethod
def _call(cls, key, patches, *args: list[Any]):
p = patches.get(key)
if p:
return p(*args)
return optimized_attention(*args[:-1], heads=args[-1]["n_heads"])
class PatchTypeModel(PatchTypeTransformer):
def set_patches(self, model_options, val):
model_options[self.name] = val[0]
def get_patches(self, model_options):
return () if self.name not in model_options else (model_options[self.name],)
class PatchTypeModelWrapper(PatchTypeModel):
@classmethod
def _call(cls, patches, apply_model, opts):
if not patches:
return apply_model(opts["input"], opts["timestep"], **opts["c"])
return patches[0](apply_model, opts)
class PatchTypeSamplerPostCfgFunction(PatchTypeModel):
def get_patches(self, model_options):
return model_options.get(self.name, ())
def set_patches(self, model_options, val):
model_options[self.name] = val
_call_result_key = "denoised"
@classmethod
def _call(cls, patches, opts):
curr_opts = opts.copy()
key = cls._call_result_key
for p in patches:
result = p(curr_opts)
curr_opts[key] = result
return result
class PatchTypeSamplerPreCfgFunction(PatchTypeSamplerPostCfgFunction):
_call_result_key = "conds_out"
class PatchTypeSamplerCfgFunction(PatchTypeModel):
@classmethod
def _call(cls, patches, opts):
if not patches:
cond_pred, uncond_pred = opts["cond_denoised"], opts["uncond_denoised"]
return uncond_pred + (cond_pred - uncond_pred) * opts["cond_scale"]
return patches[0](opts)
PATCH_TYPES = {
"input_block_patch": PatchTypeTransformer("input_block_patch"),
"input_block_patch_after_skip": PatchTypeTransformer(
"input_block_patch_after_skip",
),
"output_block_patch": PatchTypeTransformer("output_block_patch", nresult=2),
"attn1_patch": PatchTypeTransformer("attn1_patch", nresult=3),
"attn1_output_patch": PatchTypeTransformer("attn1_output_patch"),
"attn2_patch": PatchTypeTransformer("attn2_patch", nresult=3),
"attn2_output_patch": PatchTypeTransformer("attn2_output_patch"),
"middle_patch": PatchTypeTransformer("middle_patch"),
"attn1": PatchTypeTransformerReplace("attn1"),
"attn2": PatchTypeTransformerReplace("attn2"),
"model_function_wrapper": PatchTypeModelWrapper("model_function_wrapper"),
"sampler_cfg_function": PatchTypeSamplerCfgFunction("sampler_cfg_function"),
"sampler_post_cfg_function": PatchTypeSamplerPostCfgFunction(
"sampler_post_cfg_function",
),
"sampler_pre_cfg_function": PatchTypeSamplerPostCfgFunction(
"sampler_pre_cfg_function",
),
}
class ModelConditionalState:
def __init__(self):
self.last_sigma = None
self.step = None
def update(self, sigma):
if self.last_sigma is None or sigma > self.last_sigma:
self.step = 0
elif sigma != self.last_sigma:
self.step += 1
self.last_sigma = sigma
return self.step
class ModelPatchConditional:
def __init__( # noqa: PLR0917
self,
model_default,
model_matched,
start_percent: float = 0.0,
end_percent: float = 1.0,
interval: int = 1,
base_on_default: bool = True,
):
self.options_default = dict(**model_default.model_options)
self.options_matched = dict(**model_matched.model_options)
self.start_percent = start_percent
self.end_percent = end_percent
self.interval = interval
self.base = model_default if base_on_default else model_matched
self.base_on_default = base_on_default
self.sigma_end = self.sigma_start = None
def lazy_calc_steps(self):
if self.sigma_start is not None:
return
model = self.base
count = 0
while hasattr(model, "model"):
count += 1
if count > 128:
raise ValueError("I can't handle these insane levels of modelception!")
model = model.model
self.sigma_start = model.model_sampling.percent_to_sigma(
self.start_percent,
)
self.sigma_end = model.model_sampling.percent_to_sigma(
self.end_percent,
)
def should_use_patched(self, state, opts):
self.lazy_calc_steps()
if "sigmas" in opts:
sigmas = opts["sigmas"]
elif "transformer_options" in opts:
sigmas = opts["transformer_options"]["sigmas"]
elif "c" in opts:
sigmas = opts["c"]["transformer_options"]["sigmas"]
elif "sigma" in opts:
sigmas = opts["sigma"]
else:
raise ValueError("Cannot determine sigma")
iv = self.interval
sigma = sigmas[0].item()
step = state.update(sigma)
matched = sigma <= self.sigma_start and sigma >= self.sigma_end
matched &= (step % iv) == 0 if iv > 0 else ((step + 1) % abs(iv)) > 0
return matched
def mk_patch_handler(self, pt, state, key=None):
def handler(*args: list[Any]):
matched = self.should_use_patched(state, args[-1])
# print(f">> {pt.name}, active: {matched}, step: {state.step}")
opts = self.options_matched if matched else self.options_default
return pt(key, opts, *args) if key else pt(opts, *args)
return handler
def patch(self):
state = ModelConditionalState()
base_patched = self.base.clone()
for pt in PATCH_TYPES.values():
if not (pt.exists(self.options_default) or pt.exists(self.options_matched)):
continue
# print(f"set patch {pt.name}")
if not isinstance(pt, PatchTypeTransformerReplace):
pt.set_patches(
base_patched.model_options,
[self.mk_patch_handler(pt, state)],
)
continue
pt.set_patches(
base_patched.model_options,
{
k: self.mk_patch_handler(pt, state, key=k)
for k in (
pt.get_patches(self.options_default).keys()
| pt.get_patches(self.options_matched).keys()
)
},
)
base_patched.model_options["disable_cfg1_optimization"] = (
self.options_default.get("disable_cfg1_optimization", False)
or self.options_matched.get("disable_cfg1_optimization", False)
)
return base_patched
class ModelPatchConditionalNode:
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
CATEGORY = "bleh/model_patches"
DESCRIPTION = "Experimental model patch that lets you control when other model patches are active."
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_default": (
"MODEL",
{
"tooltip": "Fallback model patches, used when start/end/interval do not match.",
},
),
"model_matched": (
"MODEL",
{"tooltip": "Model patches used when start/end/interval match."},
),
"start_percent": (
"FLOAT",
{
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.001,
"tooltip": "Start time as sampling percentage (not percentage of steps). Percentages are inclusive.",
},
),
"end_percent": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 1.0,
"step": 0.001,
"tooltip": "End time as sampling percentage (not percentage of steps). Percentages are inclusive.",
},
),
"interval": (
"INT",
{
"default": 1,
"min": -999,
"max": 999,
"tooltip": "Step interval to use model_matched. If positive 3 would mean activate every third step, if negative -3 would mean skip every third step.",
},
),
"base_on_default": (
"BOOLEAN",
{
"default": True,
"tooltip": "When true, the active set of patches will be applied to model_default, otherwise they will be applied to model_matched.",
},
),
},
}
@classmethod
def patch(
cls,
*,
model_default,
model_matched=None,
start_percent: float = 0.0,
end_percent: float = 1.0,
interval: int = 1,
base_on_default: bool = True,
):
if not model_matched or start_percent >= 1.0 or interval == 0:
return (model_default.clone(),)
mopts = getattr(model_default, "model_options", None)
if mopts is None or not isinstance(mopts, dict):
# Not an instance of ModelPatcher, apparently so we can't do anything here.
return (model_default.clone(),)
return (
ModelPatchConditional(
model_default,
model_matched,
start_percent,
end_percent,
interval,
base_on_default=base_on_default,
).patch(),
)