Various improvements to BlehBlockOps

This commit is contained in:
blepping
2024-04-23 15:04:57 -06:00
parent b35ab93f8b
commit b03e824842
2 changed files with 636 additions and 290 deletions
+68 -41
View File
@@ -210,99 +210,126 @@ Example:
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
to do with skip connections but I don't know the specifics. It's important for FreeU.
Default values are show in parenthesis next to the operation argument name. You may supply an incomplete argument list,
in which case default values will be used for the remaining arguments. Ex: `[flip]` is the same as `[flip, h]`. You may
also specify the arguments as a map, keys that aren't included will use the default values. Ex: `[flip, {direction: h}]`
**`slice`**: Applies a filtering operation on a slice of the target.
1. scale: Slice scale, `1.0` would mean apply to 100% of the target, `0.5` would mean 50% of it.
2. strength: Scales the target. `1.0` would mean 100%.
3. blend ratio: Ratio of the transformed value to blend in. `1.0` means replace it with no blending.
4. blend mode: See the blend mode section.
5. hidden mean: No idea what this does really, but FreeU V2 uses it when slicing and V1 doesn't.
1. `scale`(`1.0`): Slice scale, `1.0` would mean apply to 100% of the target, `0.5` would mean 50% of it.
2. `strength`(`1.0`): Scales the target. `1.0` would mean 100%.
3. `blend`(`1.0`): Ratio of the transformed value to blend in. `1.0` means replace it with no blending.
4. `blend_mode`(`bislerp`): See the blend mode section.
5. `use_hidden_mean`(`true`): No idea what this does really, but FreeU V2 uses it when slicing and V1 doesn't.
**`ffilter`**: Applies a Fourier filter operation to the target.
1. scale: Scales the target. `1.0` would mean 100%.
2. filter: May be a string with a predefined filter name or a list of lists defining filters. See the filter section.
3. strength: Strength of the filter. `1.0` would mean to apply it at 100%.
4. threshold: Threshold for the Fourier filter. This generally should be 1.
1. `scale`(`1.0`): Scales the target. `1.0` would mean 100%.
2. `filter`(`none`): May be a string with a predefined filter name or a list of lists defining filters. See the filter section.
3. `filter_strength`(`0.5`): Strength of the filter. `1.0` would mean to apply it at 100%.
4. `threshold`(`1`): Threshold for the Fourier filter. This generally should be 1.
**`scale_torch`**: Scales the target up or down, using PyTorch's `interpolate` function.
1. type: One of `bicubic`, `nearest`, `bilinear` or `area`.
2. scale width: Ratio to scale the width. `2.0` would mean double it, `0.5` would mean half of it.
3. scale height: As above.
4. antialias: `true` to apply antialiasing after scaling or `false`.
1. `type`(`bicubic`): One of `bicubic`, `nearest`, `bilinear` or `area`.
2. `scale_width`(`1.0`): Ratio to scale the width. `2.0` would mean double it, `0.5` would mean half of it.
3. `scale_height`(`1.0`): As above.
4. `antialias`(`false`): `true` to apply antialiasing after scaling or `false`.
**`unscale_torch`**: Scale the target to be the same size as `hsp`. Only can be used when the target isn't `hsp` and condition `type` is `output`.
Can be used to reverse a `scale` or `scale_torch` operation without having to worry about calculating the ratios to get the original size back.
1. type: Same as `scale_torch`.
2. antialias: Same as `scale_torch`.
1. `type`(`bicubic`): Same as `scale_torch`.
2. `antialias`(`false`): Same as `scale_torch`.
**`scale`**: Scales the target up or down using various functions. See the scaling functions section.
1. type width: Scaling function to use for width. Note if the type is one of the ones from `scale_torch` it cannot be combined with other scaling functions.
2. type height: As above.
3. scale width: Ratio to scale the width. `2.0` would mean double it, `0.5` would mean half of it.
4. scale height: As above.
5. antialias size: Size of the antialias kernel. Between 1 and 7 inclusive. Higher numbers seem to increase blurriness.
1. `type_width`(`bicubic`): Scaling function to use for width. Note if the type is one of the ones from `scale_torch` it cannot be combined with other scaling functions.
2. `type_height`(`bicubic`): As above.
3. `scale_width`(`1.0`): Ratio to scale the width. `2.0` would mean double it, `0.5` would mean half of it.
4. `scale_height`(`1.0`): As above.
5. `antialias_size`(`0`): Size of the antialias kernel. Between 1 and 7 inclusive. Higher numbers seem to increase blurriness.
**`unscale`**: Like `unscale_torch` except it supports more scale functions and can specify width/height scale function independently.
Same restriction as `scale`.
1. type width: Scaling function to use for width. Note if the type is one of the ones from `scale_torch` it cannot be combined with other scaling functions.
2. type height: As above.
3. antialias size: Size of the antialias kernel. Between 1 and 7 inclusive. Higher numbers seem to increase blurriness.
1. `type_width`(`bicubic`): Scaling function to use for width. Note if the type is one of the ones from `scale_torch` it cannot be combined with other scaling functions.
2. `type_height`(`bicubic`): As above.
3. `antialias_size`(`0`): Size of the antialias kernel. Between 1 and 7 inclusive. Higher numbers seem to increase blurriness.
**`flip`**: Flips the target.
1. direction: `h` for horizontal flip, `v` for vertical. Note that latents generally don't tolerate being flipped very well.
1. `direction`(`h`): `h` for horizontal flip, `v` for vertical. Note that latents generally don't tolerate being flipped very well.
**`rot90`**: Does a 90 degree rotation of the target.
1. count: Number of times to rotate (can also be negative). Note that if you rotate in a way that makes the tensors not match then stuff will probably break.
1. `count`(`1`): Number of times to rotate (can also be negative). Note that if you rotate in a way that makes the tensors not match then stuff will probably break.
also as with `flip` it generally is pretty destructive to latents.
**`roll`**: Rotates the values in a dimension of the target.
1. direction: `horizontal`, `vertical`, `channels`. Note that when `type` is `input`, `input_after_skip`, `middle` or `output` you aren't actually dealing
1. `direction`(`c`): `horizontal`, `vertical`, `channels`. Note that when `type` is `input`, `input_after_skip`, `middle` or `output` you aren't actually dealing
with a latent. The second dimension ("channels") is actually the features in the layer. Rotating them can produce some pretty weird effects.
2. amount: If it's a number greater than `-1.0` and less than `1.0` this will rotate forward or backward by a percentage of the size. Otherwise it is
2. `amount`(`1`): If it's a number greater than `-1.0` and less than `1.0` this will rotate forward or backward by a percentage of the size. Otherwise it is
interpreted as the number of items to rotate forward or backward.
**`roll_channels`**: Same as `roll` but you only specify the count, it always targets channels and you can't use percentages.
1. count: Number of channels to rotate. May be negative.
1. `count`(`1`): Number of channels to rotate. May be negative.
**`target_skip`**: Changes the target.
1. If `true` will target `hsp`, otherwise will target `h`. Targeting `hsp` is only allowed when `type` is `output`, no effect otherwise.
1. `active`(`true`): 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.
1. value: Multiplier. `2.0` would double all values in the target.
1. `factor`(`1.0`): Multiplier. `2.0` would double all values in the target.
**`antialias`**: Applies an antialias effect to the target. Works the same ase with `scale`.
1. antialias size: A number between 1 and 7.
1. `size`(`7`): The antialias kernel size as a number between 1 and 7.
**`noise`**: Adds noise to the target. Can only be used when sigmas are connected. Noise will be scaled by `sigma - sigma_next`.
1. scale: Additionally scale the noise by the supplied factor. `1.0` would mean no scaling, `2.0` would double it, etc.
1. `scale`(`0.5`): Additionally scale the noise by the supplied factor. `1.0` would mean no scaling, `2.0` would double it, etc.
2. `type`(`gaussian`): Only `gaussian` unless [ComfyUI-sonar](https://github.com/blepping/ComfyUI-sonar) is installed and active, otherwise
you may use the additional noise types Sonar provides.
3. `scale_mode`(`sigdiff`): `sigdiff` scales the noise by the current sigma minus the next (requires sigmas connected),
`sigma` scales by the current sigma, `none` or an invalid type uses no scaling (you get exactly `noise * scale`).
**`debug`**: Outputs some debug information about the state.
**`blend_op`**: Allows applying a blend function to the result of another operation.
1. blend ratio: Ratio of the transformed value to blend in.
2. blend mode: See the blend mode section.
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.
1. `blend`(`1.0`): Ratio of the transformed value to blend in.
2. `blend_mode`(`bislerp`): See the blend mode section.
3. `ops`(empty): 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.
**`pad`**: Pads the target.
1. `mode`(`reflect`): One of `constant`, `reflect`, `replicate`, `circular` - see https://pytorch.org/docs/stable/generated/torch.nn.functional.pad.html#torch.nn.functional.pad
2. `top`(`0`): Amount of top padding. If this is a floating point value, it will be treated as a percentage of the dimension.
3. `bottom`(`0`): " " "
4. `left`(`0`): " " "
5. `right`(`0`): " " "
6. `constant`(`0`): Constant value to use, only applies when mode is `constant`.
_Note_: If you pad `input` (rather than `input_after_skip`) then you will need to crop the corresponding block in `output`
for both `h` and `hsp` (i.e. with `target_skip`).
**`crop`**: Crops the target.
1. `top`(`0`): Items to crop from the top. If this is a floating point value, it will be treated as a percentage of the dimension.
2. `bottom`(`0`): " " "
3. `left`(`0`): " " "
4. `right`(`0`): " " "
**`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`.
1. `scale_mode`(`bicubic`) type: Same as with `scale`.
2. `antialias`(`7`) size: Same as with `scale`.
3. `mask`(mask targeting corners): A two dimensional list of mask values. See below.
4. `ops`(empty): Same as with `blend_op`.
Simple example of a mask:
@@ -320,8 +347,8 @@ 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`.
* If a column item is a list, the first element is interpreted as the repeat count and the remaining items are repeated
however many times. Ex: `[2, 1.2, 0.5]` as a column would expand to `1.2, 0.5, 1.2, 0.5`.
These two shortcuts can be combined. A mask of `[["rep", 2, 1, [3, 0], 2]]` expands to:
+568 -249
View File
@@ -1,7 +1,10 @@
# Adapted from the ComfyUI built-in node
from __future__ import annotations
import bisect
import importlib
import operator as pyop
from collections import OrderedDict
from enum import Enum, auto
import torch
@@ -9,6 +12,54 @@ import yaml
from ..latent_utils import * # noqa: TID252
try:
sonar_noise = importlib.import_module("custom_nodes.ComfyUI-sonar.py.noise")
get_noise_sampler = sonar_noise.get_noise_sampler
except ImportError:
def get_noise_sampler(noise_type, x, *_args: list, **_kwargs: dict):
if noise_type != "gaussian":
raise ValueError("Only gaussian noise supported")
return lambda _s, _sn: torch.randn_like(x)
class CondType(Enum):
TYPE = auto()
BLOCK = auto()
STAGE = auto()
FROM_PERCENT = auto()
TO_PERCENT = auto()
PERCENT = 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 PatchType(Enum):
LATENT = auto()
INPUT = auto()
INPUT_AFTER_SKIP = auto()
MIDDLE = auto()
OUTPUT = auto()
POST_CFG = auto()
PRE_APPLY_MODEL = auto()
POST_APPLY_MODEL = auto()
class CompareType(Enum):
EQ = auto()
NE = auto()
GT = auto()
LT = auto()
GE = auto()
LE = auto()
NOT = auto()
OR = auto()
AND = auto()
class OpType(Enum):
# scale, strength, blend, blend mode, use hidden mean
@@ -56,43 +107,93 @@ class OpType(Enum):
# size
ANTIALIAS = auto()
# strength
# scale, type, scale_mode (none, sigma, sigdiff)
NOISE = auto()
# none
DEBUG = auto()
# mode (constant, reflect, replicate, circular), top, bottom, left, right, constant
PAD = auto()
class CondType(Enum):
TYPE = auto()
BLOCK = auto()
STAGE = auto()
FROM_PERCENT = auto()
TO_PERCENT = auto()
PERCENT = 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()
# top, bottom, left, right
CROP = auto()
# count, [ops]
REPEAT = auto()
class CompareType(Enum):
EQ = auto()
NE = auto()
GT = auto()
LT = auto()
GE = auto()
LE = auto()
NOT = auto()
OR = auto()
AND = auto()
OP_DEFAULTS = {
OpType.SLICE: OrderedDict(
scale=1.0,
strength=1.0,
blend=1.0,
blend_mode="bislerp",
use_hidden_mean=True,
),
OpType.FFILTER: OrderedDict(
scale=1.0,
filter="none",
filter_strength=0.5,
threshold=1,
),
OpType.SCALE_TORCH: OrderedDict(
type="bicubic",
scale_width=1.0,
scale_height=None,
antialias=False,
),
OpType.SCALE: OrderedDict(
type_width="bicubic",
type_height="bicubic",
scale_width=1.0,
scale_height=None,
antialias_size=0,
),
OpType.UNSCALE_TORCH: OrderedDict(
type="bicubic",
antialias=False,
),
OpType.UNSCALE: OrderedDict(
type_width="bicubic",
type_height="bicubic",
antialias_size=0,
),
OpType.FLIP: OrderedDict(direction="h"),
OpType.ROT90: OrderedDict(count=1),
OpType.ROLL_CHANNELS: OrderedDict(count=1),
OpType.ROLL: OrderedDict(direction="c", amount=1),
OpType.TARGET_SKIP: OrderedDict(active=True),
OpType.MULTIPLY: OrderedDict(factor=1.0),
OpType.BLEND_OP: OrderedDict(blend=1.0, blend_mode="bislerp", ops=()),
OpType.MASK_EXAMPLE_OP: OrderedDict(
scale_mode="bicubic",
antialias_size=7,
mask=(
(0.5, 0.25, (16, 0.0), 0.25, 0.5),
("rep", 18, (20, 0.0)),
(0.5, 0.25, (16, 0.0), 0.25, 0.5),
),
ops=(),
),
OpType.ANTIALIAS: OrderedDict(size=7),
OpType.NOISE: OrderedDict(scale=0.5, type="gaussian", scale_mode="sigdiff"),
OpType.DEBUG: OrderedDict(),
OpType.PAD: OrderedDict(
mode="reflect",
top=0,
bottom=0,
left=0,
right=0,
constant=None,
),
OpType.CROP: OrderedDict(top=0, bottom=0, left=0, right=0),
OpType.REPEAT: OrderedDict(count=2, ops=()),
}
class Compare:
VALID_TYPES = { # noqa: RUF012
CondType.TYPE,
CondType.BLOCK,
CondType.STAGE,
CondType.PERCENT,
@@ -132,7 +233,11 @@ class Compare:
class Condition:
def __init__(self, typ: str, value):
self.typ = getattr(CondType, typ.upper().strip())
if self.typ is not CondType.COND:
if self.typ == CondType.TYPE:
if not isinstance(value, (list, tuple)):
value = (value,)
self.value = {getattr(PatchType, pt.strip().upper()) for pt in value}
elif 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:])
@@ -178,6 +283,9 @@ class ConditionGroup:
def test(self, state: dict) -> bool:
return all(c.test(state) for c in self.conds)
def get_all_types(self) -> set[str]:
pass
def __repr__(self) -> str:
return f"<ConditionGroup: {self.conds}>"
@@ -196,217 +304,383 @@ def hidden_mean(h):
class Operation:
IDX = 0
def __init__(self, typ: str, *args: list):
self.typ = getattr(OpType, typ.upper().strip())
self.args = args
def __init__(self, typ: str | OpType, *args: list):
if isinstance(typ, str):
typ = getattr(OpType, typ.upper().strip())
self.typ = typ
defaults = OP_DEFAULTS[self.typ]
if len(args) == 1 and isinstance(args[0], dict):
args = args[0]
extra = set(args.keys()) - set(defaults.keys())
if extra:
errstr = f"Unexpected argument keys for operation {typ}: {extra}"
raise ValueError(errstr)
self.args = tuple(args.get(k, v) for k, v in defaults.items())
else:
if len(args) > len(defaults):
raise ValueError("Too many arguments for operation")
self.args = (*args, *tuple(defaults.values())[len(args) :])
@staticmethod
def build(typ: str | OpType, *args: list) -> object:
if isinstance(typ, str):
typ = getattr(OpType, typ.upper().strip())
return OP_TO_OPCLASS[typ](typ, *args)
def eval(self, state: dict):
t = out = state[state["target"]]
match self.typ:
case OpType.SCALE_TORCH | OpType.UNSCALE_TORCH:
if self.typ == OpType.SCALE_TORCH:
mode, scale_w, scale_h, antialias = self.args
width, height = (
round(t.shape[-1] * scale_w),
round(t.shape[-2] * scale_h),
)
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
width, height = hsp.shape[-1], hsp.shape[-2]
out = scale_samples(
t,
width,
height,
mode,
antialias_size=8 if antialias else 0,
)
case OpType.SCALE | OpType.UNSCALE:
if self.typ == OpType.SCALE:
mode_w, mode_h, scale_w, scale_h, antialias_size = self.args
width, height = (
round(t.shape[-1] * scale_w),
round(t.shape[-2] * scale_h),
)
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
width, height = hsp.shape[-1], hsp.shape[-2]
out = scale_samples(
t,
width,
height,
mode=mode_w,
mode_h=mode_h,
antialias_size=antialias_size,
)
case OpType.FLIP:
out = torch.flip(
t,
dims=(2 if self.args[0] == "v" else 3,),
)
case OpType.ROT90:
out = torch.rot90(t, self.args[0], dims=(3, 2))
case OpType.ROLL_CHANNELS:
out = torch.roll(t, self.args[0], dims=(1,))
case OpType.ROLL:
dims, amount = self.args
if isinstance(dims, str):
match dims:
case "h" | "horizontal":
dims = (3,)
case "v" | "vertical":
dims = (2,)
case "c" | "channels":
dims = (1,)
case _:
raise ValueError("Bad roll direction")
elif isinstance(dims, int):
dims = (dims,)
if isinstance(amount, float) and amount < 1.0 and amount > -1.0:
if len(dims) > 1:
raise ValueError(
"Cannot use percentage based amount with multiple roll dimensions",
)
amount = int(t.shape[dims[0]] * amount)
out = torch.roll(t, amount, dims=dims)
case OpType.TARGET_SKIP:
if state.get("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:
scale, filt, strength, threshold = self.args
if isinstance(filt, str):
filt = FILTER_PRESETS[filt]
out = ffilter(t, threshold, scale, filt, strength)
case OpType.SLICE:
scale, strength, blend, mode, use_hm = self.args
slice_size = round(t.shape[1] * scale)
sliced = t[:, :slice_size]
if use_hm:
result = sliced * ((strength - 1) * hidden_mean(t) + 1)
else:
result = sliced * strength
if blend != 1:
result = BLENDING_MODES[mode](sliced, result, blend)
out[:, :slice_size] = result
case OpType.MULTIPLY:
out *= self.args[0]
case OpType.BLEND_OP:
blend, mode, subops = self.args
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 = 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:
# mask = torch.ones(t.shape[2:], device=t.device, dtype=t.dtype)
# ms = 32
# mask[ms:-ms, :] = 0
# mask[:, ms:-ms] = 0
noise = torch.randn_like(t) # * mask
step_scale = state["sigma"] - state["sigma_next"]
t += noise * step_scale * self.args[0]
case OpType.DEBUG:
stcopy = {
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}")
case _:
raise ValueError("Unhandled")
out = self.op(state[state["target"]], state)
state[state["target"]] = out
def __repr__(self) -> str:
return f"<Operation({self.typ}): {self.args!r}>"
class SubOpsOperation(Operation):
SUBOPS_IDXS = ()
def __init__(self, *args: list, **kwargs: dict):
super().__init__(*args, **kwargs)
for argidx in self.SUBOPS_IDXS:
subops = self.args[argidx]
if subops and isinstance(subops[0], str):
# Simple single subop.
subops = (subops,)
compiled_subops = []
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.build(subop[0], *subop[1:])
compiled_subops.append(subop)
temp = list(self.args)
temp[argidx] = compiled_subops
self.args = tuple(temp)
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]
if use_hm:
result = sliced * ((strength - 1) * hidden_mean(t) + 1)
else:
result = sliced * strength
if blend != 1:
result = BLENDING_MODES[mode](sliced, result, blend)
out[:, :slice_size] = result
return out
class OpFFilter(Operation):
def op(self, t, _state):
scale, filt, strength, threshold = self.args
if isinstance(filt, str):
filt = FILTER_PRESETS[filt]
elif filt is None:
filt = ()
return ffilter(t, threshold, scale, filt, strength)
class OpScaleTorch(Operation):
def op(self, t, state):
if self.typ == OpType.SCALE_TORCH:
mode, scale_w, scale_h, antialias = self.args
width, height = (
round(t.shape[-1] * scale_w),
round(t.shape[-2] * scale_h),
)
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 t
mode, antialias = self.args
width, height = hsp.shape[-1], hsp.shape[-2]
return scale_samples(
t,
width,
height,
mode,
antialias_size=8 if antialias else 0,
)
class OpUnscaleTorch(OpScaleTorch):
pass
class OpScale(Operation):
def op(self, t, state):
if self.typ == OpType.SCALE:
mode_w, mode_h, scale_w, scale_h, antialias_size = self.args
width, height = (
round(t.shape[-1] * scale_w),
round(t.shape[-2] * scale_h),
)
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 t
mode_w, mode_h, antialias_size = self.args
width, height = hsp.shape[-1], hsp.shape[-2]
return scale_samples(
t,
width,
height,
mode=mode_w,
mode_h=mode_h,
antialias_size=antialias_size,
)
class OpUnscale(OpScale):
pass
class OpFlip(Operation):
def op(self, t, _state):
return torch.flip(
t,
dims=(2 if self.args[0] in ("v", "vertical") else 3,),
)
class OpRot90(Operation):
def op(self, t, _state):
return torch.rot90(t, self.args[0], dims=(3, 2))
class OpRollChannels(Operation):
def op(self, t, _state):
return torch.roll(t, self.args[0], dims=(1,))
class OpRoll(Operation):
def op(self, t, _state):
dims, amount = self.args
if isinstance(dims, str):
match dims:
case "h" | "horizontal":
dims = (3,)
case "v" | "vertical":
dims = (2,)
case "c" | "channels":
dims = (1,)
case _:
raise ValueError("Bad roll direction")
elif isinstance(dims, int):
dims = (dims,)
if isinstance(amount, float) and amount < 1.0 and amount > -1.0:
if len(dims) > 1:
raise ValueError(
"Cannot use percentage based amount with multiple roll dimensions",
)
amount = int(t.shape[dims[0]] * amount)
return torch.roll(t, amount, dims=dims)
class OpTargetSkip(Operation):
def op(self, t, state):
if state.get("hsp") is None:
if state["target"] == "hsp":
state["target"] = "h"
return t
state["target"] = "hsp" if self.args[0] is True else "h"
return t
class OpMultiply(Operation):
def op(self, t, _state):
return t.mul_(selg.args[0])
class OpBlendOp(SubOpsOperation):
SUBOPS_IDXS = (2,)
def op(self, t, state):
blend, mode, subops = self.args
tempname = f"temp{Operation.IDX}"
Operation.IDX += 1
old_target = state["target"]
state[tempname] = t.clone()
for subop in subops:
state["target"] = tempname
subop.eval(state)
state["target"] = old_target
out = BLENDING_MODES[mode](t, state[tempname], blend)
del state[tempname]
return out
class OpMaskExampleOp(SubOpsOperation):
SUBOPS_IDXS = (3,)
def __init__(self, *args: list, **kwargs: dict):
super().__init__(*args, **kwargs)
scale_mode, antialias_size, maskdef, subops = self.args
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=torch.float32, device="cpu")
self.args = (scale_mode, antialias_size, mask, subops)
def op(self, t, state):
scale_mode, antialias_size, mask, subops = self.args
mask = scale_samples(
mask.view(1, 1, *mask.shape).to(t.device, dtype=t.dtype),
t.shape[-1],
t.shape[-2],
mode=scale_mode,
antialias_size=antialias_size,
).broadcast_to(t.shape)
tempname = f"temp{Operation.IDX}"
Operation.IDX += 1
old_target = state["target"]
state[tempname] = t.clone()
for subop in subops:
state["target"] = tempname
subop.eval(state)
state["target"] = old_target
out = state[tempname] * mask
out += t * (1 - mask)
del state[tempname]
return out
class OpAntialias(Operation):
def op(self, t, _state):
return antialias_tensor(t, self.args[0])
class OpNoise(Operation):
def op(self, t, state):
scale, noise_type, scale_mode = self.args
match scale_mode:
case "sigma":
step_scale = state.get("sigma", 1.0)
case "sigdiff":
if "sigma" in state and "sigma_next" in state:
step_scale = state["sigma"] - state["sigma_next"]
else:
step_scale = state.get("sigma", 1.0)
case _:
step_scale = 1.0
noise_sampler = get_noise_sampler(
noise_type,
t,
state["sigma_min"],
state["sigma_max"],
)
noise = noise_sampler(state.get("sigma"), state.get("sigma_next"))
t += noise * step_scale * scale
return t
class OpDebug(Operation):
def op(self, t, state):
stcopy = {
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}")
return t
class OpPad(Operation):
def op(self, t, _state):
mode, top, bottom, left, right, constant_value = self.args
if mode != "constant":
constant_value = None
shp = t.shape
top, bottom = tuple(
val if isinstance(val, int) else int(shp[-2] * val) for val in (top, bottom)
)
left, right = tuple(
val if isinstance(val, int) else int(shp[-1] * val) for val in (left, right)
)
return torch.nn.functional.pad(
t,
(left, right, top, bottom),
mode=mode,
value=constant_value,
)
class OpCrop(Operation):
def op(self, t, _state):
top, bottom, left, right = self.args
shp = t.shape
top, bottom = tuple(
val if isinstance(val, int) else int(shp[-2] * val) for val in (top, bottom)
)
left, right = tuple(
val if isinstance(val, int) else int(shp[-1] * val) for val in (left, right)
)
bottom, right = shp[-2] - bottom, shp[-1] - right
return t[:, :, top:bottom, left:right]
class OpRepeat(SubOpsOperation):
SUBOPS_IDXS = (1,)
def op(self, _t, state):
count, subops = self.args
for _ in range(count):
for subop in subops:
subop.eval(state)
return state[state["target"]]
OP_TO_OPCLASS = {
OpType.SLICE: OpSlice,
OpType.FFILTER: OpFFilter,
OpType.SCALE_TORCH: OpScaleTorch,
OpType.UNSCALE_TORCH: OpUnscaleTorch,
OpType.SCALE: OpScale,
OpType.UNSCALE: OpUnscale,
OpType.FLIP: OpFlip,
OpType.ROT90: OpRot90,
OpType.ROLL_CHANNELS: OpRollChannels,
OpType.ROLL: OpRoll,
OpType.TARGET_SKIP: OpTargetSkip,
OpType.MULTIPLY: OpMultiply,
OpType.BLEND_OP: OpBlendOp,
OpType.MASK_EXAMPLE_OP: OpMaskExampleOp,
OpType.ANTIALIAS: OpAntialias,
OpType.NOISE: OpNoise,
OpType.DEBUG: OpDebug,
OpType.PAD: OpPad,
OpType.CROP: OpCrop,
OpType.REPEAT: OpRepeat,
}
class Rule:
@classmethod
def from_dict(cls, val) -> object:
@@ -421,13 +695,14 @@ class Rule:
nomatched=d.get("else", ()),
)
for d in val
if not d.get("disable")
)
def __init__(self, conds=(), ops=(), matched=(), 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.ops = tuple(Operation.build(o[0], *o[1:]) for o in ops)
self.matched = Rule.from_dict(matched)
self.nomatched = Rule.from_dict(nomatched)
@@ -441,11 +716,11 @@ class Rule:
def eval(self, state: dict) -> None:
# print("EVAL", state | {"h": None, "hsp": None})
if not self.conds.test(state):
for r in self.nomatched:
r.eval(state)
return
state["target"] = "h"
for o in self.ops:
o.eval(state)
for r in self.matched:
@@ -459,13 +734,15 @@ class RuleGroup:
@classmethod
def from_yaml(cls, s: str) -> object:
parsed_rules = yaml.safe_load(s)
return cls(tuple(Rule.from_dict(r)[0] for r in parsed_rules))
return cls(tuple(r for rs in parsed_rules for r in Rule.from_dict(rs)))
def __init__(self, rules):
self.rules = rules
def eval(self, state):
def eval(self, state, toplevel=False):
for rule in self.rules:
if toplevel:
state["target"] = "h"
rule.eval(state)
return state
@@ -505,8 +782,9 @@ class BlehBlockOps:
# Arbitrary number that should have good enough precision
pct_steps = 400
pct_incr = 1.0 / pct_steps
model_sampling = model.get_model_object("model_sampling")
sig2pct = tuple(
model.model.model_sampling.percent_to_sigma(x / pct_steps)
model_sampling.percent_to_sigma(x / pct_steps)
for x in range(pct_steps, -1, -1)
)
@@ -520,22 +798,28 @@ class BlehBlockOps:
return pct_incr * (pct_steps - idx)
def set_state_step(state, sigma):
sdict = {
CondType.STEP: -1,
CondType.STEP_EXACT: -1,
"sigma": sigma,
"sigma_min": model_sampling.sigma_min,
"sigma_max": model_sampling.sigma_max,
}
if sigmas_opt is None:
state[CondType.STEP_EXACT] = state[CondType.STEP] = -1
state |= sdict
return state
sigmadiff, idx = torch.min(torch.abs(sigmas_opt[:-1] - sigma), 0)
idx = idx.item()
state |= {
state |= sdict | {
CondType.STEP: idx + 1,
CondType.STEP_EXACT: -1 if sigmadiff.item() > 1.5e-06 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: str, topts: dict, h, hsp=None):
def make_state(typ: PatchType, topts: dict, h, hsp=None):
pct = get_pct(topts)
if pct is None:
return None
@@ -558,13 +842,13 @@ class BlehBlockOps:
state = make_state(typ, topts, h)
if state is None:
return h
return rules.eval(state)["h"]
return rules.eval(state, toplevel=True)["h"]
def output_block_patch(h, hsp, transformer_options: dict):
state = make_state("output", transformer_options, h, hsp)
state = make_state(PatchType.OUTPUT, transformer_options, h, hsp)
if state is None:
return h
rules.eval(state)
rules.eval(state, toplevel=True)
return state["h"], state["hsp"]
def post_cfg_patch(args: dict):
@@ -572,7 +856,7 @@ class BlehBlockOps:
if pct is None:
return None
state = {
CondType.TYPE: "post_cfg",
CondType.TYPE: PatchType.POST_CFG,
CondType.PERCENT: pct,
CondType.BLOCK: -1,
CondType.STAGE: -1,
@@ -585,11 +869,11 @@ class BlehBlockOps:
m = model.clone()
m.set_model_input_block_patch_after_skip(
lambda *args: block_patch("input_after_skip", *args),
lambda *args: block_patch(PatchType.INPUT_AFTER_SKIP, *args),
)
m.set_model_input_block_patch(lambda *args: block_patch("input", *args))
m.set_model_input_block_patch(lambda *args: block_patch(PatchType.INPUT, *args))
m.set_model_patch(
lambda *args: block_patch("middle", *args),
lambda *args: block_patch(PatchType.MIDDLE, *args),
"middle_block_patch",
)
m.set_model_output_block_patch(output_block_patch)
@@ -597,6 +881,41 @@ class BlehBlockOps:
post_cfg_patch,
disable_cfg1_optimization=True,
)
orig_model_function_wrapper = model.model_options.get("model_function_wrapper")
def pre_model(state):
state[CondType.TYPE] = PatchType.PRE_APPLY_MODEL
return rules.eval(state, toplevel=True)["h"]
def post_model(state, result):
state[CondType.TYPE] = PatchType.POST_APPLY_MODEL
state["target"] = "h"
state["h"] = result
return rules.eval(state, toplevel=True)["h"]
def model_unet_function_wrapper(apply_model, args):
pct = get_pct({"sigmas": args["timestep"]})
if pct is None:
return None
state = {
CondType.PERCENT: pct,
CondType.BLOCK: -1,
CondType.STAGE: -1,
"h": args["input"],
"hsp": None,
"target": "h",
}
set_state_step(state, args["timestep"].max().item())
x = pre_model(state)
args = args | {"input": x}
if orig_model_function_wrapper is not None:
result = orig_model_function_wrapper(apply_model, args)
else:
result = apply_model(args["input"], args["timestep"], **args["c"])
return post_model(state, result)
m.set_model_unet_function_wrapper(model_unet_function_wrapper)
return (m,)
@@ -662,11 +981,11 @@ class BlehLatentOps:
}
RETURN_TYPES = ("LATENT",)
FUNCTION = "upscale"
FUNCTION = "go"
CATEGORY = "latent"
def upscale(
def go(
self,
samples,
rules: str,
@@ -678,7 +997,7 @@ class BlehLatentOps:
rules = RuleGroup.from_yaml(rules)
stensor = samples["samples"]
state = {
CondType.TYPE: "latent",
CondType.TYPE: PatchType.LATENT,
CondType.PERCENT: 0.0,
CondType.BLOCK: -1,
CondType.STAGE: -1,
@@ -686,6 +1005,6 @@ class BlehLatentOps:
"hsp": None,
"target": "h",
}
rules.eval(state)
rules.eval(state, toplevel=True)
samples["samples"] = state["h"]
return (samples,)