From f0721e09c9986ffc878ce7ab9eae264705486ded Mon Sep 17 00:00:00 2001 From: Daniel Martinek Date: Sun, 5 Apr 2026 00:33:35 +0200 Subject: [PATCH] README closer to rality --- README.md | 32 ++++++++------- more_math/AudioMathNode.py | 13 +++--- more_math/Parser/MathExpr.g4 | 6 +-- more_math/Parser/UnifiedMathVisitor.py | 22 ++++++---- more_math/SelectiveGuiderMathNode.py | 56 +++++++++++++++++--------- more_math/nodes.py | 2 + 6 files changed, 82 insertions(+), 49 deletions(-) diff --git a/README.md b/README.md index 0ad3cc5..f42c7a3 100644 --- a/README.md +++ b/README.md @@ -177,8 +177,8 @@ You can also get the node from comfy manager under the name of More math. ### Image -- `blur(x, sigma[,auto_convert])` or `gaussian`: Applies a Gaussian blur with given `sigma` along last two dimensions (auto_convert is 0 or omitted) or attempts to use correct spatial dimensions if auto_convert is 1. -- `edge(x,[kernel_size[,auto_convert]]`: Applies a Sobel edge detection filter along the last two dimension or spatial dimensions (Height and Width) - can be selected by optional value (0 or missing = use last 2 dimensions). +- `blur(x, sigma)` or `gaussian`: Applies a Gaussian blur with given `sigma` along last two dimensions (auto_convert is 0 or omitted) or attempts to use correct spatial dimensions if auto_convert is 1. +- `edge(x,[kernel_size]`: Applies a Sobel edge detection filter along the last two dimension or spatial dimensions (Height and Width) - can be selected by optional value (0 or missing = use last 2 dimensions). - `ezconvolution(tensor, kw, [kh], [kd], k_expr)` or `ezconv`: Applies a convolution to `tensor`. Automatically permutes tensor to try to make it work with various inputs without the need to permute manually. - `k_expr` can be a math expression (using `kX`, `kY`, `kZ`) or a list literal. - `convolution(tensor, kw, [kh], [kd], k_expr)` or `conv`: Applies a convolution to `tensor`. Does not perform automatic permutations. Expects layout `(Batch, Channel, Spatial...)`. @@ -378,12 +378,11 @@ Adds support for hooking into specific layers/blocks of the model during guided Current implementation details: -- `hook_target` is currently hardcoded to `all` in code. -- `hook_when` is currently set but not used for filtering. +- `hook_target` supports runtime filtering in node UI. - `layer_x` is used as direct index match (`idx == layer_x`) for current hook context. - For guiders with `original_conds`, base conditions are restored before reattaching hooks to avoid hook accumulation across reruns/interrupted runs. - Active hook paths currently include: - - Attention override (`attn1` / `attn2` / `attn_unknown`) + - Attention override (`attn1` / `attn2` / `double_block_attn` / `single_block_attn` / `attn_unknown`) - DiT block replace (`dit.double_block`, `dit.single_block`) - UNet block patches (`input_block_patch`, `middle_patch`, `output_block_patch`) - Timestep embedding start via `emb_patch` (`block_name="time_emb"`, `layer_x=0`) @@ -415,12 +414,14 @@ The following variables are available in `Expression`. | `F` | list/tensor | Collection of all float inputs. | | `D0..Dn` | tensor | Per-dimension index tensors from `generate_dim_variables`. | | `S0..Sn` | float | Per-dimension sizes from `generate_dim_variables`. | -| `hook_kind` | string | Active hook identifier: `attn1`, `attn2`, `attn_unknown`, `dit_block`, `unet_block`, `model_begin`, `model_end`, or `unknown`. | +| `hook_kind` | string | Active hook identifier: `attn1`, `attn2`, `double_block_attn`, `single_block_attn`, `attn_unknown`, `dit_block`, `unet_block`, `model_begin`, `model_end`, or `unknown`. | | `hook_domain` | string | High-level domain: `attention`, `diffusion`, or `unknown`. | -| `attn_kind` | string | Attention kind: `attn1`, `attn2`, `attn_unknown`, or `none` outside attention hooks. | +| `attn_kind` | string | Attention kind: `attn1`, `attn2`, `double_block_attn`, `single_block_attn`, `attn_unknown`, or `none` outside attention hooks. | | `transformer_index` | float | Attention sub-block index inside a UNet block (`-1` if unavailable). | | `is_attn1` | float (0/1) | `1` when current hook is `attn1`, else `0`. | | `is_attn2` | float (0/1) | `1` when current hook is `attn2`, else `0`. | +| `is_attn1_hook` | float (0/1) | `1` when `attn_kind=="attn1"`, else `0`. | +| `is_attn2_hook` | float (0/1) | `1` when `attn_kind=="attn2"`, else `0`. | | `is_dit` | float (0/1) | `1` when current hook is DiT block hook, else `0`. | | `is_unet_block` | float (0/1) | `1` when current hook is UNet block hook, else `0`. | | `is_time_emb` | float (0/1) | `1` when current hook is timestep embedding entry (`block_name=="time_emb"`), else `0`. | @@ -436,30 +437,33 @@ The following variables are available in `Expression`. | `v` | tensor | Value tensor in attention hooks; fallback placeholder otherwise. | | `heads` | float | Number of attention heads (attention hooks only, else `0`). | | `dim_head` | float | Per-head channel size (`q.shape[-1] / heads`) when available. | -| `activations_shape` | list | Raw shape from transformer context (typically `[B, C, H, W]` in UNet attention). Empty list if unavailable. | +| `activations_shape` | list | Raw shape from transformer context. Empty list if unavailable. | | `activation_b` | float | Batch dimension from `activations_shape[0]` (or `-1`). | | `activation_c` | float | Channel dimension from `activations_shape[1]` (or `-1`). | | `activation_h` | float | Height dimension from `activations_shape[2]` (or `-1`). | | `activation_w` | float | Width dimension from `activations_shape[3]` (or `-1`). | -| `attn_mode` | string | Unified attention mode: `self`, `cross`, or `unknown`. Prefer this over `attn1`/`attn2` for non-UNet models. | +| `attn_mode` | string | Legacy compatibility field (default `unknown`). | +| `attention_relation` | string | Inferred semantic relation: `self`, `cross`, or `unknown`. | | `is_self_attention` | float (0/1) | `1` when the active attention is self-attention. | | `is_cross_attention` | float (0/1) | `1` when the active attention is cross-attention. | | `has_context` | float (0/1) | `1` when attention context appears to be present. | | `query_tokens` | float | Query sequence length. | | `context_tokens` | float | Context sequence length. | | `value_tokens` | float | Value sequence length. | -| `activation_rank` | float | Rank of `activations_shape` (`4` for image-like, `5` for video-like tensors). | +| `activation_rank` | float | Rank of `activations_shape`. | | `activation_t` | float | Temporal dimension for video-like activations (`-1` if unavailable). | #### Practical notes - On SD1.x, repeated hits on the same `layer_id` are normal in attention because one UNet block can contain multiple transformer sub-blocks. -- Use `transformer_index` to target exactly one sub-block, for example: - `(attn_kind=="attn2" * transformer_index==0) ? : inp` +- Use `transformer_index` to target exactly one sub-block. - For timestep-begin hooking use `layer_x=0` and filter by `block_name=="time_emb"` (or `layer_key=="unet.time_emb.0"`). - For model-edge hooks filter by `hook_kind=="model_begin"` or `hook_kind=="model_end"`. - For model-agnostic expressions, prefer guard variables: `has_qkv`, `is_dit`, `is_unet_block`, `is_attn1`, `is_attn2`. - - `attn2` is only a hook label, not guaranteed to mean real cross-attention. - Use `is_cross_attention` only as an inferred relation from runtime metadata. -- Use `is_attn2_hook` when you need to know the hook path, and `attention_relation` when you need the semantic relation. +- Use `attention_relation` for semantic relation (`self`/`cross`) and `attn_kind` for hook-path classification. + +- Selective guider math: + - `hook_target`: `all`, `dit_block`, `unet_block`, `attn1`, `attn2`, + `double_block_attn`, `single_block_attn`, `model_begin`, `model_end` diff --git a/more_math/AudioMathNode.py b/more_math/AudioMathNode.py index e37e795..34f3f59 100644 --- a/more_math/AudioMathNode.py +++ b/more_math/AudioMathNode.py @@ -21,9 +21,9 @@ class AudioMathNode(io.ComfyNode): Enables math expressions on Audio. Inputs: - I: Autogrow image inputs (I0, I1, ...) + V: Autogrow audio inputs (V0, V1, ...) F: Autogrow float inputs (F0, F1, ...) - Image: Expression + Expression: Expression """ @classmethod @@ -45,7 +45,7 @@ class AudioMathNode(io.ComfyNode): options=["do nothing","error","tile", "pad"], display_name="on size mismatch", default="error", - tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero." + tooltip="How to handle mismatched audio shapes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing samples as zero." ), io.Int.Input(id="batching", default=0), io.Boolean.Input( @@ -108,14 +108,17 @@ class AudioMathNode(io.ComfyNode): "y": F.get("F2", 0.0) if F.get("F2") is not None else 0.0, "z": F.get("F3", 0.0) if F.get("F3") is not None else 0.0, "B": getIndexTensorAlongDim(a_w, 0), + "batch": getIndexTensorAlongDim(a_w, 0), "C": getIndexTensorAlongDim(a_w, 1), "channel": getIndexTensorAlongDim(a_w, 1), + "N": float(a_w.shape[1]), + "channel_count": float(a_w.shape[1]), "S": getIndexTensorAlongDim(a_w, 2), "sample": getIndexTensorAlongDim(a_w, 2), + "T": float(a_w.shape[2]), + "sample_count": float(a_w.shape[2]), "R": sample_rate, "sample_rate": sample_rate, - "batch": getIndexTensorAlongDim(a_w, 0), - "T": float(a_w.shape[0]), "batch_count": float(a_w.shape[0]), } | generate_dim_variables(a_w) | V_norm_waveforms | sample_rates diff --git a/more_math/Parser/MathExpr.g4 b/more_math/Parser/MathExpr.g4 index a918f8a..8709d95 100644 --- a/more_math/Parser/MathExpr.g4 +++ b/more_math/Parser/MathExpr.g4 @@ -150,8 +150,8 @@ func1: | ARGSORT LPAREN expr (COMMA expr)? RPAREN # ArgsortFunc | ARGMIN LPAREN expr RPAREN # ArgminFunc | ARGMAX LPAREN expr RPAREN # ArgmaxFunc - | SOFTMAX LPAREN expr RPAREN # SoftmaxFunc - | SOFTMIN LPAREN expr RPAREN # SoftminFunc + | SOFTMAX LPAREN expr (COMMA expr)? RPAREN # SoftmaxFunc + | SOFTMIN LPAREN expr (COMMA expr)? RPAREN # SoftminFunc | ERF LPAREN expr RPAREN # ErfFunc | ERFINV LPAREN expr RPAREN # ErfinvFunc | UNIQUE LPAREN expr RPAREN # UniqueFunc @@ -365,7 +365,7 @@ FLATTEN: 'flatten'; APPEND: 'append'; GET_VALUE: 'get_value'; FLOW_APPLY: 'flow_apply'; -BATCH_SHUFFLE: 'batch_shuffle' | 'shuffle'; +BATCH_SHUFFLE: 'batch_shuffle' | 'shuffle' | 'select'; MOTION_MASK: 'motion_mask'; FLOW_TO_IMAGE: 'flow_to_image'; OVERLAY: 'overlay'; diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index dff8ce2..268e710 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -2389,11 +2389,20 @@ class UnifiedMathVisitor(MathExprVisitor): return torch.full(shape, value, device=self.device,dtype=type) def visitSoftmaxFunc(self, ctx): - val = self._promote_to_tensor((yield ctx.expr())) - return F.softmax(val.float()) + val = self._promote_to_tensor((yield ctx.expr(0))).float() + dim = -1 + if len(ctx.expr()) > 1: + dim_val = (yield ctx.expr(1)) + dim = self._to_int(dim_val, ctx, "softmax dim") + return F.softmax(val, dim=dim) + def visitSoftminFunc(self, ctx): - val = self._promote_to_tensor((yield ctx.expr())) - return F.softmax(-val.float()) + val = self._promote_to_tensor((yield ctx.expr(0))).float() + dim = -1 + if len(ctx.expr()) > 1: + dim_val = (yield ctx.expr(1)) + dim = self._to_int(dim_val, ctx, "softmin dim") + return F.softmax(-val, dim=dim) def visitArgminFunc(self, ctx): val = self._promote_to_tensor((yield ctx.expr())) @@ -3360,7 +3369,6 @@ class UnifiedMathVisitor(MathExprVisitor): return torch.stack([r, g, b], dim=-1) - def visitEntropyFunc(self, ctx): val = self._promote_to_tensor((yield ctx.expr())) # Shannon entropy: -sum(p * log(p)) @@ -3520,6 +3528,4 @@ class UnifiedMathVisitor(MathExprVisitor): def visitErfinvFunc(self, ctx): """erfinv(x) - inverse error function""" x = self._promote_to_tensor((yield ctx.expr())) - return torch.erfinv(x) - - \ No newline at end of file + return torch.erfinv(x) \ No newline at end of file diff --git a/more_math/SelectiveGuiderMathNode.py b/more_math/SelectiveGuiderMathNode.py index 86214c9..f132e36 100644 --- a/more_math/SelectiveGuiderMathNode.py +++ b/more_math/SelectiveGuiderMathNode.py @@ -78,10 +78,18 @@ class SelectiveGuiderMathNode(io.ComfyNode): tree = parse_expr(Expression) if isinstance(Expression, str) else Expression MAX_HOOK_INDEX = 999 - hit_flags = {"attn1": False, "attn2": False, "dit": False, "attn_unknown": False, "unet": False} + hit_flags = { + "attn1": False, + "attn2": False, + "double_block_attn": False, + "single_block_attn": False, + "dit": False, + "attn_unknown": False, + "unet": False, + } def resolve_attn_kind(transformer_options: dict, q: torch.Tensor | None = None, k: torch.Tensor | None = None, v: torch.Tensor | None = None) -> str: - # Detekce Flux / MM-DiT architektury (vracíme specifické názvosloví) + # Flux / MM-DiT btype = transformer_options.get("block_type", None) if btype == "double": return "double_block_attn" @@ -99,7 +107,6 @@ class SelectiveGuiderMathNode(io.ComfyNode): ), ) - # Detekce z block_name pro non-flux modely blk = transformer_options.get("block", None) if isinstance(blk, (tuple, list)) and len(blk) > 0: blk_name = str(blk[0]).lower() @@ -122,13 +129,15 @@ class SelectiveGuiderMathNode(io.ComfyNode): if int(raw) == 1: return "attn2" - # Fallback heuristika - if isinstance(q, torch.Tensor) and isinstance(k, torch.Tensor) and isinstance(v, torch.Tensor): - if (q.data_ptr() == k.data_ptr()) and (k.data_ptr() == v.data_ptr()): - return "attn1" - if q.data_ptr() != k.data_ptr(): + # Bezpečný fallback podle délek tokenů (funguje i pro FLOW self-attn) + if isinstance(q, torch.Tensor) and isinstance(k, torch.Tensor): + if q.ndim >= 2 and k.ndim >= 2: + q_tokens = int(q.shape[-2]) + k_tokens = int(k.shape[-2]) + if q_tokens == k_tokens: + return "attn1" return "attn2" - + return "attn_unknown" def resolve_x(total_blocks: int | None) -> int: @@ -205,6 +214,8 @@ class SelectiveGuiderMathNode(io.ComfyNode): "is_negative": 0.0, "is_attn1_hook": 0.0, "is_attn2_hook": 0.0, + "is_double_block_attn_hook": 0.0, + "is_single_block_attn_hook": 0.0, "attention_relation": "unknown", } | generate_dim_variables(inp) @@ -237,6 +248,10 @@ class SelectiveGuiderMathNode(io.ComfyNode): variables["is_attn1_hook"] = 1.0 elif variables.get("attn_kind") == "attn2": variables["is_attn2_hook"] = 1.0 + elif variables.get("attn_kind") == "double_block_attn": + variables["is_double_block_attn_hook"] = 1.0 + elif variables.get("attn_kind") == "single_block_attn": + variables["is_single_block_attn_hook"] = 1.0 if hasattr(qv, "shape") and len(qv.shape) >= 2: variables["query_tokens"] = float(qv.shape[-2]) @@ -327,19 +342,25 @@ class SelectiveGuiderMathNode(io.ComfyNode): def attn_override(original_attn, q, k, v, heads, **kwargs): transformer_options = kwargs.get("transformer_options", {}) attn_kind = resolve_attn_kind(transformer_options, q=q, k=k, v=v) - t_index = int(transformer_options.get("transformer_index", -1)) - - if hook_target != "all" and hook_target in ("attn1", "attn2", "double_block_attn", "single_block_attn"): - if hook_target != attn_kind: - return original_attn(q, k, v, heads, **kwargs) blk = transformer_options.get("block", None) + btype = str(transformer_options.get("block_type", "unknown")) + bindex = int(transformer_options.get("block_index", -1)) + if isinstance(blk, (tuple, list)) and len(blk) >= 2: stage = str(blk[0]) layer_id = int(blk[1]) else: - stage = "unknown" - layer_id = int(transformer_options.get("block_index", -1)) + # lepší fallback pro Flux/FLOW + stage = btype if btype in ("double", "single") else str(transformer_options.get("model_type", "unknown")).lower() + layer_id = bindex + + t_index_raw = transformer_options.get("transformer_index", None) + t_index = int(t_index_raw) if isinstance(t_index_raw, (int, float)) else layer_id + + if hook_target != "all" and hook_target in ("attn1", "attn2", "double_block_attn", "single_block_attn"): + if hook_target != attn_kind: + return original_attn(q, k, v, heads, **kwargs) idx = layer_id total_blocks = None @@ -360,9 +381,6 @@ class SelectiveGuiderMathNode(io.ComfyNode): if not isinstance(act_shape, (list, tuple)): act_shape = [] - act_h = float(act_shape[2]) if len(act_shape) >= 4 else -1.0 - act_w = float(act_shape[3]) if len(act_shape) >= 4 else -1.0 - meta = { "hook_kind": attn_kind, "hook_domain": "attention", diff --git a/more_math/nodes.py b/more_math/nodes.py index 7c023a6..6686912 100644 --- a/more_math/nodes.py +++ b/more_math/nodes.py @@ -26,6 +26,7 @@ from .AudioToSpectrogramNode import AudioToSpectrogram from .NoiseMathNode import NoiseMathNode from .AudioMathNode import AudioMathNode +from .StringMathNode import StringMathNode from .SelectiveGuiderMathNode import SelectiveGuiderMathNode from .ScriptTextWindow import ScriptTextInput @@ -169,6 +170,7 @@ class MoreMathExtension(ComfyExtension): GuiderMathNode, NoiseMathNode, AudioMathNode, + StringMathNode, ScriptTextInput, SelectiveGuiderMathNode ]