From 235703abb4fdb318388e4090070bc6a50182383e Mon Sep 17 00:00:00 2001 From: blepping Date: Tue, 14 Jan 2025 06:45:21 -0700 Subject: [PATCH] Added BlehLatentBlend node. Added BlehCast node. Added BlehSetSigmas node. BlockOps improvements. Other misc cleanups and refactoring. --- README.md | 15 ++++ __init__.py | 3 + changelog.md | 7 ++ py/nodes/misc.py | 171 ++++++++++++++++++++++++++++++++++++++++++- py/nodes/ops.py | 108 +++++++++++++++++++++------ py/nodes/samplers.py | 2 +- 6 files changed, 280 insertions(+), 26 deletions(-) diff --git a/README.md b/README.md index 090b4b3..0e4c068 100644 --- a/README.md +++ b/README.md @@ -231,6 +231,21 @@ Like the builtin `LatentScaleBy` node, however it allows setting the horizontal as well as allowing providing an extended list of scaling options. Can also be useful for testing what different types of scaling or enhancement effects look like. +### BlehLatentBlend + +Allows blending latents using any of the blending modes available. + +### BlehCast + +Advanced node: Allows tricking ComfyUI into thinking a value of one type is a different type. This does not actually convert anything, just lets you connect things that otherwise couldn't be connected. In other words, don't do it unless you know the actual object is compatible with the input. + +### BlehSetSigmas + +Advanced sigma manipulation node which can be used to insert sigmas into other sigmas, adjust them, replace them or +just manually enter a list of sigmas. Note: Experimental, not well tested. + +*** + ## Scaling Types * bicubic: Generally the safe option. diff --git a/__init__.py b/__init__.py index e0fda1f..8918e05 100644 --- a/__init__.py +++ b/__init__.py @@ -33,11 +33,14 @@ NODE_CLASS_MAPPINGS = { "BlehInsaneChainSampler": samplers.BlehInsaneChainSampler, "BlehLatentOps": ops.BlehLatentOps, "BlehLatentScaleBy": ops.BlehLatentScaleBy, + "BlehLatentBlend": ops.BlehLatentBlend, "BlehModelPatchConditional": modelPatchConditional.ModelPatchConditionalNode, "BlehPlug": misc.BlehPlug, "BlehRefinerAfter": refinerAfter.BlehRefinerAfter, "BlehSageAttentionSampler": sageAttention.BlehSageAttentionSampler, "BlehSetSamplerPreset": samplers.BlehSetSamplerPreset, + "BlehCast": misc.BlehCast, + "BlehSetSigmas": misc.BlehSetSigmas, } NODE_DISPLAY_NAME_MAPPINGS = { diff --git a/changelog.md b/changelog.md index ca7bf7a..6eaf013 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,13 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. +## 20250114 + +* Added `BlehLatentBlend` node. +* Added `BlehCast` node that lets crazy people connect things that shouldn't be connected. +* Added `BlehSetSigmas` node. +* Some BlockOps functions have expanded capabilities now. + ## 20250109 * The strategy the SageAttention nodes use to patch ComfyUI's attention should work for third-party custom nodes more reliably now. Please create an issue if you experience problems. diff --git a/py/nodes/misc.py b/py/nodes/misc.py index 418587e..f0c7a02 100644 --- a/py/nodes/misc.py +++ b/py/nodes/misc.py @@ -1,6 +1,8 @@ from __future__ import annotations +import operator import random +from decimal import Decimal import torch from comfy import model_management @@ -87,7 +89,7 @@ class BlehDisableNoise: def go( cls, noise_seed: int, - seed_offset: None | int = 1, + seed_offset: int | None = 1, ) -> tuple[SeededDisableNoise]: return ( SeededDisableNoise( @@ -97,7 +99,7 @@ class BlehDisableNoise: ) -class Wildcard(str): +class Wildcard(str): # noqa: FURB189 __slots__ = () def __ne__(self, _unused): @@ -120,3 +122,168 @@ class BlehPlug: @classmethod def go(cls): return (None,) + + +class BlehCast: + DESCRIPTION = "UNSAFE: This node allows casting its input to any type. NOTE: This does not actually change the data in any way, it just allows you to connect its output to any input. Only use if you know for sure the data is compatible." + FUNCTION = "go" + CATEGORY = "hacks" + + WILDCARD = Wildcard("*") + RETURN_TYPES = (WILDCARD,) + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "any_input": ( + cls.WILDCARD, + { + "forceInput": True, + "description": "You can connect any type of input here, but take to ensure that you connect the output from this node to an input that is compatible.", + }, + ), + }, + } + + @classmethod + def go(cls, *, any_input): + return (any_input,) + + +class BlehSetSigmas: + DESCRIPTION = "Advanced node that allows manipulating SIGMAS. For example, you can manually enter a list of sigmas, insert some new sigmas into existing SIGMAS, etc." + FUNCTION = "go" + CATEGORY = "sampling/custom_sampling/sigmas" + RETURN_TYPES = ("SIGMAS",) + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "start_index": ( + "INT", + { + "default": 0, + "tooltip": "Start index for modifying sigmas, zero-based. May be set to a negative value to index from the end, i.e. -1 is the last item, -2 is the penultimate item.", + }, + ), + "mode": ( + ("replace", "insert", "multiply", "add", "subtract", "divide"), + { + "default": "replace", + "tooltip": "", + }, + ), + "order": ( + ("AB", "BA"), + { + "default": "AB", + "tooltip": "Only applies to add, subtract, multiply and divide operations. Controls the order of operations. For example if order AB then add means A*B, if order BA then add means B*A.", + }, + ), + "commasep_sigmas_b": ( + "STRING", + { + "default": "", + "tooltip": "Exclusive with sigmas_b. Enter a comma-separated list of sigma values here. For non-insert mode, the input sigmas will be padded with zeros if necessary. Example: start_index=2 (3rd item), mode=replace, input sigmas 4,3,2,1 and you used replace mode with 0.3,0.2,0.1 the output would be 4,3,0.3,0.2,0.1", + }, + ), + }, + "optional": { + "sigmas_a": ( + "SIGMAS", + { + "forceInput": True, + "tooltip": "Optional input as long as commasep_sigmas is not also empty. If not supplied, an initial sigmas list of the appropriate size will be generated filled with zeros.", + }, + ), + "sigmas_b": ( + "SIGMAS", + { + "forceInput": True, + "tooltip": "Optionally populate this or commasep_sigmas_b but not both.", + }, + ), + }, + } + + OP_MAP = { # noqa: RUF012 + "add": operator.add, + "subtract": operator.sub, + "multiply": operator.mul, + "divide": operator.truediv, + } + + @classmethod + def go( + cls, + *, + start_index: int, + mode: str, + order: str, + commasep_sigmas_b: str, + sigmas_a: torch.Tensor | None = None, + sigmas_b: torch.Tensor | None = None, + ) -> tuple: + new_sigmas_list = tuple( + Decimal(val) for val in commasep_sigmas_b.strip().split(",") if val.strip() + ) + if new_sigmas_list and sigmas_b is not None: + raise ValueError( + "Must populate one of sigmas_b or commasep_sigmas_b but not both.", + ) + if sigmas_b is not None: + sigmas_b = sigmas_b.to(dtype=torch.float64, device="cpu", copy=True) + else: + sigmas_b = torch.tensor(new_sigmas_list, device="cpu", dtype=torch.float64) + newlen = sigmas_b.numel() + if sigmas_a is None or sigmas_a.numel() == 0: + sigmas_a = None + if not newlen: + raise ValueError( + "sigmas_a, commasep_sigmas_b and sigmas_b can't all be empty.", + ) + if start_index < 0: + raise ValueError( + "Negative start_index doesn't make sense when input sigmas are empty.", + ) + if newlen == 0: + return (sigmas_a.to(dtype=torch.float, copy=True),) + oldlen = 0 if sigmas_a is None else sigmas_a.numel() + if start_index < 0: + start_index = oldlen + start_index + if start_index < 0: + raise ValueError( + "Negative start index points past the beginning of sigmas_a", + ) + past_end = 0 if start_index < oldlen else start_index + 1 - oldlen + if past_end and mode == "insert": + mode = "replace" + if past_end: + outlen = oldlen + newlen + past_end - 1 + elif mode == "insert": + outlen = oldlen + newlen + else: + outlen = oldlen + max(0, newlen - (oldlen - start_index)) + sigmas_out = torch.zeros(outlen, device="cpu", dtype=torch.float64) + if mode == "insert": + sigmas_out[:start_index] = sigmas_a[:start_index] + sigmas_out[start_index : start_index + newlen] = sigmas_b + sigmas_out[start_index + newlen :] = sigmas_a[start_index:] + else: + if oldlen: + sigmas_out[:oldlen] = sigmas_a + if mode == "replace": + sigmas_out[start_index : start_index + newlen] = sigmas_b + else: + opfun = cls.OP_MAP.get(mode) + if opfun is None: + raise ValueError("Bad mode") + arga = sigmas_out[start_index : start_index + newlen] + if order == "BA": + arga, argb = sigmas_b, arga + else: + argb = sigmas_b + sigmas_out[start_index : start_index + newlen] = opfun(arga, argb) + return (sigmas_out.to(torch.float),) diff --git a/py/nodes/ops.py b/py/nodes/ops.py index 0a609cd..515f056 100644 --- a/py/nodes/ops.py +++ b/py/nodes/ops.py @@ -63,7 +63,7 @@ class CompareType(Enum): class OpType(Enum): - # scale, strength, blend, blend mode, use hidden mean + # scale, strength, blend, blend mode, use hidden mean, dim, scale offset SLICE = auto() # scale, filter, filter strength, threshold @@ -102,7 +102,7 @@ class OpType(Enum): # blend strength, blend_mode, [op] BLEND_OP = auto() - # scale mode, antialias size, mask example, [op] + # scale mode, antialias size, mask example, [op], blend_mode MASK_EXAMPLE_OP = auto() # size @@ -134,6 +134,8 @@ OP_DEFAULTS = { blend=1.0, blend_mode="bislerp", use_hidden_mean=True, + dim=1, + scale_offset=0, ), OpType.FFILTER: OrderedDict( scale=1.0, @@ -179,6 +181,7 @@ OP_DEFAULTS = { (0.5, 0.25, (16, 0.0), 0.25, 0.5), ), ops=(), + blend_mode="lerp", ), OpType.ANTIALIAS: OrderedDict(size=7), OpType.NOISE: OrderedDict(scale=0.5, type="gaussian", scale_mode="sigdiff"), @@ -368,16 +371,26 @@ class SubOpsOperation(Operation): class OpSlice(Operation): def op(self, t, _state): out = t - scale, strength, blend, mode, use_hm = self.args - slice_size = round(t.shape[1] * scale) - sliced = t[:, :slice_size] + scale, strength, blend, mode, use_hm, dim, scale_offset = self.args + if dim < 0: + dim = t.ndim + dim + dim_size = t.shape[dim] + slice_size = max(1, round(dim_size * scale)) + slice_offset = int(dim_size * scale_offset) + slice_def = tuple( + slice(None, None) + if idx != dim + else slice(slice_offset, slice_offset + slice_size) + for idx in range(dim + 1) + ) + sliced = t[slice_def] if use_hm: - result = sliced * ((strength - 1) * hidden_mean(t) + 1) + result = sliced * ((strength - 1) * hidden_mean(t)[slice_def] + 1) else: result = sliced * strength if blend != 1: result = BLENDING_MODES[mode](sliced, result, blend) - out[:, :slice_size] = result + out[slice_def] = result return out @@ -457,10 +470,14 @@ class OpUnscale(OpScale): class OpFlip(Operation): def op(self, t, _state): - return torch.flip( - t, - dims=(2 if self.args[0] in {"v", "vertical"} else 3,), - ) + dimarg = self.args[0] + if isinstance(dimarg, str): + dim = dimarg[:1] == "v" + elif isinstance(dimarg, int): + dim = (dimarg,) + else: + dim = dimarg + return torch.flip(t, dims=dim) class OpRot90(Operation): @@ -535,7 +552,10 @@ class OpMaskExampleOp(SubOpsOperation): def __init__(self, *args: list, **kwargs: dict): super().__init__(*args, **kwargs) - scale_mode, antialias_size, maskdef, subops = self.args + scale_mode, antialias_size, maskdef, subops, blend_mode = self.args + blend_function = BLENDING_MODES.get(blend_mode) + if blend_function is None: + raise ValueError("Bad blend mode") mask = [] for rowidx in range(len(maskdef)): repeats = 1 @@ -551,10 +571,16 @@ class OpMaskExampleOp(SubOpsOperation): row.append(col) mask += (row,) * repeats mask = torch.tensor(mask, dtype=torch.float32, device="cpu") - self.args = (scale_mode, antialias_size, mask, subops) + self.args = ( + scale_mode, + antialias_size, + mask, + subops, + blend_function, + ) def op(self, t, state): - scale_mode, antialias_size, mask, subops = self.args + scale_mode, antialias_size, mask, subops, blend_function = self.args mask = scale_samples( mask.view(1, 1, *mask.shape).to(t.device, dtype=t.dtype), t.shape[-1], @@ -570,8 +596,7 @@ class OpMaskExampleOp(SubOpsOperation): state["target"] = tempname subop.eval(state) state["target"] = old_target - out = state[tempname] * mask - out += t * (1 - mask) + out = blend_function(t, state[tempname], mask) del state[tempname] return out @@ -748,6 +773,8 @@ class RuleGroup: @classmethod def from_yaml(cls, s: str) -> object: parsed_rules = yaml.safe_load(s) + if parsed_rules is None: + return cls(()) return cls(tuple(r for rs in parsed_rules for r in Rule.from_dict(rs))) def __init__(self, rules): @@ -786,7 +813,7 @@ class BlehBlockOps: cls, model, rules: str, - sigmas_opt: None | torch.Tensor = None, + sigmas_opt: torch.Tensor | None = None, ): rules = rules.strip() if len(rules) == 0: @@ -980,8 +1007,8 @@ class BlehLatentScaleBy: stensor, width, height, - method_horizontal, - method_vertical, + mode=method_horizontal, + mode_h=method_vertical, antialias_size=antialias_size, ) return (samples,) @@ -995,6 +1022,9 @@ class BlehLatentOps: "samples": ("LATENT",), "rules": ("STRING", {"multiline": True, "dynamicPrompts": False}), }, + "optional": { + "samples_hsp": ("LATENT",), + }, } RETURN_TYPES = ("LATENT",) @@ -1005,8 +1035,10 @@ class BlehLatentOps: @classmethod def go( cls, - samples, + *, + samples: dict, rules: str, + samples_hsp: dict | None = None, ): samples = samples.copy() rules = rules.strip() @@ -1020,9 +1052,39 @@ class BlehLatentOps: CondType.BLOCK: -1, CondType.STAGE: -1, "h": stensor, - "hsp": None, + "hsp": None if samples_hsp is None else samples_hsp["samples"], "target": "h", } rules.eval(state, toplevel=True) - samples["samples"] = state["h"] - return (samples,) + return ({"samples": state["h"]},) + + +class BlehLatentBlend: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "samples1": ("LATENT",), + "samples2": ("LATENT",), + "samples2_percent": ("FLOAT", {"default": 0.5}), + "blend_mode": (tuple(BLENDING_MODES.keys()),), + }, + } + + RETURN_TYPES = ("LATENT",) + FUNCTION = "go" + + CATEGORY = "latent" + + @classmethod + def go( + cls, + *, + samples1: dict, + samples2: dict, + samples2_percent=0.5, + blend_mode="lerp", + ): + a, b = samples1["samples"], samples2["samples"] + blend_function = BLENDING_MODES[blend_mode] + return ({"samples": blend_function(a, b, samples2_percent)},) diff --git a/py/nodes/samplers.py b/py/nodes/samplers.py index ec279e8..f0cb8b5 100644 --- a/py/nodes/samplers.py +++ b/py/nodes/samplers.py @@ -21,7 +21,7 @@ with contextlib.suppress(Exception): BLEH_PRESET_LIMIT, max( 0, - int(environ.get("COMFYUI_BLEH_SAMPLER_PRESET_COUNT", 1)), + int(environ.get("COMFYUI_BLEH_SAMPLER_PRESET_COUNT", "1")), ), )