BlehBlockOps: Add generic logic/comparison condition type.

BlehBlockOps: Add mask_example_op operation.

Other BlehBlockOps improvements/cleanups.
This commit is contained in:
blepping
2024-04-12 11:02:17 -06:00
parent 50b0ac7665
commit 0386f8802a
3 changed files with 277 additions and 60 deletions
+75 -5
View File
@@ -121,7 +121,7 @@ SDXL is not going to work at all.
Very experimental advanced node that allows defining model patches using YAML. This node is still under development and may be changed.
**Note**: ComfyUI seems to strip out curly braces so you can't use YAML's inline object notation.
**Note**: ComfyUI seems to strip out curly braces so you can't use YAML's inline object notation. You can use `<` and `>` instead.
The top level YAML should consist of a list of objects with a condition `if`, a list of `ops` that run if the condition succeeds.
Objects `then` and `else` also take the same form as the top level object and apply when the `if` condition matches (or not in the case of `else`).
@@ -149,21 +149,29 @@ All object fields (`if`, `then`, `else`, `ops`) are optional. An empty object is
#### Conditions
**`type`**: One of `input`, `input_after_skip`, `middle`, `output`, `latent`, `post_cfg`.
**`type`**: One of `input`, `input_after_skip`, `middle`, `output` (preceding are block patches), `latent`, `post_cfg`.
**Note**: ComfyUI doesn't allow patching the middle blocks by default, this feature is only available if you have
[FreeU Advanced](https://github.com/WASasquatch/FreeU_Advanced) installed and enabled. (It patches ComfyUI to support patching
the middle blocks.)
**`block`**: The block number. Only applies when type is `input`, `input_after_skip`, `middle` or `output`.
**`stage`**: The model stage. Applies to the same types as `block`. You can think of this in terms of FreeU's `b1`, `b2` - the number is the stage.
**`from_percent`**: Matches when sampling is greater or equal to the percent. Note that this is sampling percentage, not percentage of steps.
**`percent`**: Percentage of sampling completed as a number between `0.0` and `1.0`. Note that this is sampling percentage, not percentage of steps.
Does not apply to type `latent`.
**`from_percent`**: Matches when sampling is greater or equal to the percent. Same restrictions as `percent`.
**`to_percent`**: Matches when sampling is less or equal to the percent. Same restrictions as `from_percent`.
**`step`**: Only applies when sigmas are connected to the `BlehBlockOps` node. A step will be determined as the index of the closest
matching sigma. In other words, if you don't connect sigmas that exactly match the sigmas used for sampling you won't get accurate steps.
Does not apply to type `latent`.
**`step_exact`**: Same restrictions as `step`, however will only be set if the current sigma _exactly_ matches a step. Otherwise the
value will be `-1`.
**`from_step`**: As above, but matches when the step is greater or equal to the value.
**`from_step`**: As above, but matches when the step is less or equal to the value.
@@ -171,6 +179,34 @@ Does not apply to type `latent`.
**`step_interval`**: Same restrictions as the other step condition types. Matches when the step modulus interval is 0. In other words,
every other step starting from the first step you'd use an interval of `2` and the `then` branch (since `1 % 2 == 1` which is not 0).
**`cond`**: Generic condition, has two forms:
*Comparison*: Takes three arguments: comparison type (`eq`, `ne`, `gt`, `lt`, `ge`, `le`), a condition type with
a numeric value (`block`, `stage`, `percent`, `step`, `step_exact`) and a value or list of values to compare with.
Example:
```yaml
- if: [cond, [lt, percent, 0.35]]
```
*Logic*: Takes a logic operation type (`not`, `and`, `or`) and a list of condition blocks. **Note**: The logic operation is applied
to the result of the condition block and not the fields within it.
Example:
```yaml
- if:
cond: [not,
[cond, [or,
[cond, [lt, step, 1]],
[cond, [gt, step, 5]],
]]
] # A verbose way of expressing step >= 1 and step <= 5
- if:
- [cond, [ge, step, 1]]
- [cond, [le, step, 5]] # Same as above
- if: [[from_step, 1], [to_step, 5]] # Also same as above
```
#### Operations
Operations mostly modify a target which can be `h` or `hsp`. `hsp` is only a valid target when `type` is `output`. I think it has something
@@ -241,7 +277,7 @@ Same restriction as `scale`.
**`target_skip`**: Changes the target.
1. If `true` will target `hsp`, otherwise will target `h`. Targeting `hsp` is only allowed when `type` is `output`.
1. If `true` will target `hsp`, otherwise will target `h`. Targeting `hsp` is only allowed when `type` is `output`, no effect otherwise.
**`multiply`**: Multiply the target by the value.
@@ -261,8 +297,42 @@ Same restriction as `scale`.
1. blend ratio: Ratio of the transformed value to blend in.
2. blend mode: See the blend mode section.
3. op: The operation as a list, with the name first. i.e. `[blend_op, 0.5, inject, [multiply, 0.5]]`
3. ops: The operation as a list, with the name first. i.e. `[blend_op, 0.5, inject, [multiply, 0.5]]`. May also be a list of operations.
**`mask_example_op`**: Applies providing a mask by example and masks the result of an operation or list of operations.
1. scale type: Same as with `scale`.
2. antialias size: Same as with `scale`.
3. mask: A two dimensional list of mask values. See below.
4. ops: Same as with `blend_op`.
Simple example of a mask:
```plaintext
[ [1.0, 0.0, 0.0, 1.0],
[0.0, 0.0, 0.0, 0.0],
[1.0, 0.0, 0.0, 1.0],
]
```
With this mask, the result of the mask ops will be applied at full strength to the corners. The mask is scaled up to
the size of the target tensor, so with this example the masked corners will be proportionately quite large if the
latent or tensor is much bigger than the mask. There are two convenience tricks for defining larger masks without
having to specify each value:
* If the first element in a row is `"rep"` then the second element is interpreted as a row repeat count and the
rest of the items in the row constitute the row. Ex: `["rep", 2, 1, 0, 1]` expands to two rows of `1, 0, 1`.
* If a column item is a list, the first element is interpreted as the repeat count and the second as the element
- the rest is ignored. Ex: `[3, 1.2]` as a column would expand to `1.2, 1.2, 1.2`.
These two shortcuts can be combined. A mask of `[["rep", 2, 1, [3, 0], 2]]` expands to:
```plaintext
[
[1, 0, 0, 0, 2],
[1, 0, 0, 0, 2],
]
```
#### Blend Modes
+4
View File
@@ -2,6 +2,10 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 20240412
* Added `BlehBlockOps` and `BlehLatentOps` nodes.
## 20240403
* Added `BlehRefinerAfter` node.
+198 -55
View File
@@ -1,6 +1,7 @@
# Adapted from the ComfyUI built-in node
import bisect
import operator as pyop
from enum import Enum, auto
import torch
@@ -49,6 +50,9 @@ class OpType(Enum):
# blend strength, blend_mode, [op]
BLEND_OP = auto()
# scale mode, antialias size, mask example, [op]
MASK_EXAMPLE_OP = auto()
# size
ANTIALIAS = auto()
@@ -66,25 +70,74 @@ class CondType(Enum):
FROM_PERCENT = auto()
TO_PERCENT = auto()
PERCENT = auto()
STEP = auto()
STEP = auto() # Calculated from closest sigma.
STEP_EXACT = auto() # Only exact matching sigma or -1.
FROM_STEP = auto()
TO_STEP = auto()
STEP_INTERVAL = auto()
COND = auto()
class BlockType(Enum):
INPUT = auto()
INPUT_AFTER_SKIP = auto()
MIDDLE = auto()
OUTPUT = auto()
class CompareType(Enum):
EQ = auto()
NE = auto()
GT = auto()
LT = auto()
GE = auto()
LE = auto()
NOT = auto()
OR = auto()
AND = auto()
class BlockCond:
def __init__(self, typ, value):
class Compare:
VALID_TYPES = { # noqa: RUF012
CondType.TYPE,
CondType.BLOCK,
CondType.STAGE,
CondType.PERCENT,
CondType.STEP,
CondType.STEP_EXACT,
}
def __init__(self, typ: str, value):
self.typ = getattr(CompareType, typ.upper().strip())
if self.typ in (CompareType.OR, CompareType.AND, CompareType.NOT):
self.value = tuple(ConditionGroup(v) for v in value)
self.field = None
return
self.field = getattr(CondType, value[0].upper().strip())
if self.field not in self.VALID_TYPES:
raise TypeError("Invalid type compare operation")
self.opfn = getattr(pyop, self.typ.name.lower())
self.value = value[1:]
if not isinstance(self.value, (list, tuple)):
self.value = (self.value,)
def test(self, state: dict) -> bool:
match self.typ:
case CompareType.NOT:
return all(not v.test(state) for v in self.value)
case CompareType.AND:
return all(v.test(state) for v in self.value)
case CompareType.OR:
return any(v.test(state) for v in self.value)
opfn, fieldval = self.opfn, state[self.field]
return all(opfn(fieldval, val) for val in self.value)
def __repr__(self) -> str:
return f"<Compare({self.typ}): {self.field}, {self.value}>"
class Condition:
def __init__(self, typ: str, value):
self.typ = getattr(CondType, typ.upper().strip())
self.value = set(value if isinstance(value, (list, tuple)) else (value,))
if self.typ is not CondType.COND:
self.value = set(value if isinstance(value, (list, tuple)) else (value,))
else:
self.value = Compare(value[0], value[1:])
def test(self, state):
def test(self, state: dict) -> bool:
match self.typ:
case CondType.FROM_PERCENT:
pct = state[CondType.PERCENT]
@@ -101,28 +154,32 @@ class BlockCond:
case CondType.STEP_INTERVAL:
step = state[CondType.STEP]
result = step > 0 and all(step % v == 0 for v in self.value)
case CondType.COND:
result = self.value.test(state)
case _:
result = state[self.typ] in self.value
return result
def __repr__(self):
def __repr__(self) -> str:
return f"<Cond({self.typ}): {self.value}>"
class BlockConds:
class ConditionGroup:
def __init__(self, conds):
if not conds:
self.conds = ()
return
if isinstance(conds, dict):
conds = tuple(conds.items())
self.conds = tuple(BlockCond(ct, cv) for ct, cv in conds)
if isinstance(conds[0], str):
conds = (conds,)
self.conds = tuple(Condition(ct, cv) for ct, cv in conds)
def test(self, state):
def test(self, state: dict) -> bool:
return all(c.test(state) for c in self.conds)
def __repr__(self):
return f"<Conds[{self.count}]: {self.conds}"
def __repr__(self) -> str:
return f"<ConditionGroup: {self.conds}>"
# Copied from https://github.com/WASasquatch/FreeU_Advanced
@@ -136,14 +193,14 @@ def hidden_mean(h):
).unsqueeze(2).unsqueeze(3)
class BlockOp:
class Operation:
IDX = 0
def __init__(self, typ, *args):
def __init__(self, typ: str, *args: list):
self.typ = getattr(OpType, typ.upper().strip())
self.args = args
def eval(self, state):
def eval(self, state: dict):
t = out = state[state["target"]]
match self.typ:
case OpType.SCALE_TORCH | OpType.UNSCALE_TORCH:
@@ -155,6 +212,10 @@ class BlockOp:
)
else:
hsp = state["hsp"]
if hsp is None:
raise ValueError(
"Can only use unscale_torch when HSP is set (output)",
)
if t.shape[-1] == hsp.shape[-1] and t.shape[-2] == hsp.shape[-2]:
return
mode, antialias = self.args
@@ -175,6 +236,10 @@ class BlockOp:
)
else:
hsp = state["hsp"]
if hsp is None:
raise ValueError(
"Can only use unscale when HSP is set (output)",
)
if t.shape[-1] == hsp.shape[-1] and t.shape[-2] == hsp.shape[-2]:
return
mode_w, mode_h, antialias_size = self.args
@@ -218,6 +283,10 @@ class BlockOp:
amount = int(t.shape[dims[0]] * amount)
out = torch.roll(t, amount, dims=dims)
case OpType.TARGET_SKIP:
if get(state, "hsp") is None:
if state["target"] == "hsp":
state["target"] = "h"
return
state["target"] = "hsp" if self.args[0] is True else "h"
return
case OpType.FFILTER:
@@ -240,15 +309,76 @@ class BlockOp:
out *= self.args[0]
case OpType.BLEND_OP:
blend, mode, subops = self.args
tempname = f"temp{BlockOp.IDX}"
BlockOp.IDX += 1
if subops and isinstance(subops[0], str):
# Simple single subop.
subops = (subops,)
tempname = f"temp{Operation.IDX}"
Operation.IDX += 1
old_target = state["target"]
state[tempname] = t.clone()
state["target"] = tempname
BlockOp(subops[0], *subops[1:]).eval(state)
for idx in range(len(subops)):
subop = subops[idx]
if isinstance(subop, dict):
# Compile to rule.
subop = subops[idx] = Rule.from_dict(subops[idx])
elif isinstance(subop, (list, tuple)):
# Compile to op.
subop = Operation(subop[0], *subop[1:])
state["target"] = tempname
subop.eval(state)
state["target"] = old_target
out = BLENDING_MODES[mode](t, state[tempname], blend)
del state[tempname]
case OpType.MASK_EXAMPLE_OP:
scale_mode, antialias_size, maskdef, subops = self.args
if not isinstance(maskdef, torch.Tensor):
# Compile the mask example.
mask = []
for rowidx in range(len(maskdef)):
repeats = 1
rowdef = maskdef[rowidx]
if rowdef and rowdef[0] == "rep":
repeats = int(rowdef[1])
rowdef = rowdef[2:]
row = []
for col in rowdef:
if isinstance(col, (list, tuple)):
row += (col[1],) * col[0]
else:
row.append(col)
mask += (row,) * repeats
mask = torch.tensor(mask, dtype=t.dtype, device="cpu")
self.args = (scale_mode, antialias_size, mask, subops)
else:
mask = maskdef
mask = scale_samples(
mask.view(1, 1, *mask.shape).to(t.device),
t.shape[-1],
t.shape[-2],
mode=scale_mode,
antialias_size=antialias_size,
).broadcast_to(t.shape)
if subops and isinstance(subops[0], str):
# Simple single subop.
subops = (subops,)
tempname = f"temp{Operation.IDX}"
Operation.IDX += 1
old_target = state["target"]
state[tempname] = t.clone()
for idx in range(len(subops)):
subop = subops[idx]
if isinstance(subop, dict):
# Compile to rule.
subop = subops[idx] = Rule.from_dict(subops[idx])
elif isinstance(subop, (list, tuple)):
# Compile to op.
subop = Operation(subop[0], *subop[1:])
state["target"] = tempname
subop.eval(state)
state["target"] = old_target
out = state[tempname] * mask
out += t * (1 - mask)
del state[tempname]
case OpType.ANTIALIAS:
out = antialias_tensor(t, self.args[0])
case OpType.NOISE:
@@ -261,7 +391,10 @@ class BlockOp:
t += noise * step_scale * self.args[0]
case OpType.DEBUG:
stcopy = {
k: v for k, v in state.items() if not isinstance(v, torch.Tensor)
k: v
if not isinstance(v, torch.Tensor)
else f"<Tensor: shape={v.shape}, dtype={v.dtype}>"
for k, v in state.items()
}
stcopy["target_shape"] = t.shape
print(f">> BlehOps debug: {stcopy!r}")
@@ -270,13 +403,13 @@ class BlockOp:
raise ValueError("Unhandled")
state[state["target"]] = out
def __repr__(self):
return f"<Op({self.typ}): {self.args!r}>"
def __repr__(self) -> str:
return f"<Operation({self.typ}): {self.args!r}>"
class BlockRule:
class Rule:
@classmethod
def from_dict(cls, val):
def from_dict(cls, val) -> object:
if not isinstance(val, (list, tuple)):
val = (val,)
@@ -291,12 +424,14 @@ class BlockRule:
)
def __init__(self, conds=(), ops=(), matched=(), nomatched=()):
self.conds = BlockConds(conds)
self.ops = tuple(BlockOp(o[0], *o[1:]) for o in ops)
self.matched = BlockRule.from_dict(matched)
self.nomatched = BlockRule.from_dict(nomatched)
self.conds = ConditionGroup(conds)
if ops and isinstance(ops[0], str):
ops = (ops,)
self.ops = tuple(Operation(o[0], *o[1:]) for o in ops)
self.matched = Rule.from_dict(matched)
self.nomatched = Rule.from_dict(nomatched)
def get_all_types(self):
def get_all_types(self) -> set:
result = {c.value for c in self.conds if c.typ == CondType.TYPE}
for r in self.matched:
result |= r.get_all_types()
@@ -304,7 +439,7 @@ class BlockRule:
result |= r.get_all_types()
return result
def eval(self, state):
def eval(self, state: dict) -> None:
# print("EVAL", state | {"h": None, "hsp": None})
if not self.conds.test(state):
for r in self.nomatched:
@@ -320,11 +455,13 @@ class BlockRule:
return f"<Rule: IF({self.conds}) THEN({self.ops}, {self.matched}) ELSE({self.nomatched})>"
class BlockRules:
class RuleGroup:
@classmethod
def from_yaml(cls, s):
def from_yaml(cls, s: str, curlybrace_hack=True) -> object:
if curlybrace_hack:
s = s.replace("<", "{").replace(">", "}")
parsed_rules = yaml.safe_load(s)
return cls(tuple(BlockRule.from_dict(r)[0] for r in parsed_rules))
return cls(tuple(Rule.from_dict(r)[0] for r in parsed_rules))
def __init__(self, rules):
self.rules = rules
@@ -334,6 +471,9 @@ class BlockRules:
rule.eval(state)
return state
def __repr__(self) -> str:
return f"<RuleGroup: {self.rules}>"
class BlehBlockOps:
RETURN_TYPES = ("MODEL",)
@@ -355,13 +495,14 @@ class BlehBlockOps:
def patch(
self,
model,
rules,
rules: str,
sigmas_opt=None,
):
rules = rules.strip()
if len(rules) == 0:
return (model.clone(),)
rules = BlockRules.from_yaml(rules)
rules = RuleGroup.from_yaml(rules)
# print("RULES", rules)
# Arbitrary number that should have good enough precision
pct_steps = 400
@@ -382,19 +523,21 @@ class BlehBlockOps:
def set_state_step(state, sigma):
if sigmas_opt is None:
state[CondType.STEP] = -1
state[Condtype.STEP_EXACT] = state[CondType.STEP] = -1
return st
step_idx = torch.argmin(torch.abs(sigmas_opt - sigma)).item()
sigmadiff, idx = torch.min(torch.abs(sigmas_opt - sigma), 0)
idx = idx.item()
state |= {
CondType.STEP: step_idx + 1,
"sigma": sigmas_opt[step_idx].item(),
"sigma_next": sigmas_opt[step_idx + 1].item(),
CondType.STEP: idx + 1,
CondType.STEP_EXACT: -1 if sigmadiff.item() != 0 else idx + 1,
"sigma": sigmas_opt[idx].item(),
"sigma_next": sigmas_opt[idx + 1].item(),
}
return state
stages = (1280, 640, 320)
def make_state(typ, topts, h, hsp=None):
def make_state(typ: str, topts: dict, h, hsp=None):
pct = get_pct(topts)
if pct is None:
return None
@@ -413,20 +556,20 @@ class BlehBlockOps:
set_state_step(result, topts["sigmas"].max().item())
return result
def block_patch(typ, h, topts):
def block_patch(typ, h, topts: dict):
state = make_state(typ, topts, h)
if state is None:
return h
return rules.eval(state)["h"]
def output_block_patch(h, hsp, transformer_options):
def output_block_patch(h, hsp, transformer_options: dict):
state = make_state("output", transformer_options, h, hsp)
if state is None:
return h
rules.eval(state)
return state["h"], state["hsp"]
def post_cfg_patch(args):
def post_cfg_patch(args: dict):
pct = get_pct({"sigmas": args["sigma"]})
if pct is None:
return None
@@ -487,11 +630,11 @@ class BlehLatentScaleBy:
def upscale(
self,
samples,
method_horizontal,
method_vertical,
scale_width,
scale_height,
antialias_size,
method_horizontal: str,
method_vertical: str,
scale_width: float,
scale_height: float,
antialias_size: int,
):
if method_vertical == "same":
method_vertical = method_horizontal
@@ -528,13 +671,13 @@ class BlehLatentOps:
def upscale(
self,
samples,
rules,
rules: str,
):
samples = samples.copy()
rules = rules.strip()
if len(rules) == 0:
return (samples,)
rules = BlockRules.from_yaml(rules)
rules = RuleGroup.from_yaml(rules)
stensor = samples["samples"]
state = {
CondType.TYPE: "latent",