From c4b270bc964c54fbaf4816468cc1fcaac320ca42 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Sun, 1 Feb 2026 16:33:33 +0100 Subject: [PATCH] AI add a variable which contains all input V variables V0 = V[0] ... --- README.md | 13 ++++++++++++ more_math/AudioMathNode.py | 7 +++++++ more_math/ConditioningMathNode.py | 14 ++++++++++++- more_math/FloatMathNode.py | 8 ++++++- more_math/GuiderMathNode.py | 7 +++++++ more_math/ImageMathNode.py | 8 ++++++- more_math/LatentMathNode.py | 9 +++++++- more_math/MaskMathNode.py | 8 ++++++- more_math/NoiseMathNode.py | 8 ++++++- more_math/VideoMathNode.py | 15 ++++++++++++- more_math/helper_functions.py | 33 +++++++++++++++++++++++++++++ more_math/modelLikeCommon.py | 8 ++++++- repro_grammar.py | 35 +++++++++++++++++++++++++++++++ 13 files changed, 165 insertions(+), 8 deletions(-) create mode 100644 repro_grammar.py diff --git a/README.md b/README.md index 19845a2..9fde168 100644 --- a/README.md +++ b/README.md @@ -23,6 +23,13 @@ You can also get the node from comfy manager under the name of More math. - Vector Math: Support for List literals `[v1, v2, ...]` and operations between lists/scalars/tensors - Custom functions `funcname(variable,variable,...)->expression;` they can be used in any later defined custom function or in expression. Shadowing inbuilt functions do not work. **Be careful with recursion. There is no stack limit. Got to 700 000 iterations before I got bored.** - Custom variables `varname=expression;` They can be used in any later assigment or final expression. +- Support for **indexed assignment**: `a[i, j, ...] = expression;`. Supports multidimensional tensors and nested lists. + - **Scalar Filling**: If the assigned value has only 1 element (scalar, 1-element list/tensor), it fills the entire selected slice. + - **Rank Matching**: Automatically squeezes leading ones from the value to match the rank of the target slice (e.g., assigning a 4D tensor with `dim0=1` to a 3D slice). +- **Available Variables**: + - `V0`, `V1`, ...: Individual input variables. + - `V`: A stacked tensor of all input variables `V` (shape: `[num_variables, ...]`). Available when shapes match. + - `Vcnt` or `V_count`: Number of input variables. - Support for control flow statements including `if/else`, `while` loops, blocks `{}`, and `return` statements. `if`/`else`/`while` do not work like ternary operator or other inbuilts. They colapse tensors and list to single value using any. - Support for stack. Stack survives between field evaluations but not between nodes or end of node execution. - Usefull in GuiderMath node to store variables between steps. @@ -37,6 +44,10 @@ You can also get the node from comfy manager under the name of More math. - Modifications to existing variables persist to outer scope - **Return Statements**: `return [expression];` - Early return from functions or top-level expressions +- **For Loops**: `for (variable in expression) statement` + - Iterates over elements of a list or a tensor (along dimension 0) +- **Break/Continue**: `break;`, `continue;` + - Control loop execution (works in `while` and `for` loops) ## Operators @@ -141,6 +152,8 @@ You can also get the node from comfy manager under the name of More math. - `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 standard PyTorch layout `(Batch, Channel, Spatial...)`. - `k_expr` can be a math expression (using `kX`, `kY`, `kZ`) or a list literal. +- **`get_value(tensor, position)`**: Retrieves a value from a tensor at the specified N-dimensional position (provided as a list or tensor). Uses the formula `pos0*strides[0] + pos1*strides[1] + ...` to find the linear index. +- **`crop(tensor, position, size)`**: Extracts a sub-tensor of specified `size` starting at `position` (both provided as lists/tensors). Areas outside the input tensor are filled with zeros. - `permute(tensor, dims)` or `perm`: Rearranges the dimensions of the tensor. (e.g., `perm(a, [2, 3, 0, 1])`) - `reshape(tensor, shape)` or `rshp`: Reshapes the tensor to a new shape. (e.g., `rshp(a, [S0*S1, S2, S3])`) diff --git a/more_math/AudioMathNode.py b/more_math/AudioMathNode.py index 524c69d..65eed6b 100644 --- a/more_math/AudioMathNode.py +++ b/more_math/AudioMathNode.py @@ -5,6 +5,7 @@ from .helper_functions import ( as_tensor, normalize_to_common_shape, make_zero_like, + get_v_variable ) from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io @@ -127,6 +128,12 @@ class AudioMathNode(io.ComfyNode): "batch_count": a_w.shape[0], } | generate_dim_variables(a_w) | V_norm_waveforms | sample_rates + v_stacked, v_cnt = get_v_variable(V_norm_waveforms, length_mismatch=length_mismatch) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + for k, val in F.items(): variables[k] = val if val is not None else 0.0 diff --git a/more_math/ConditioningMathNode.py b/more_math/ConditioningMathNode.py index db752c7..85ecf6f 100644 --- a/more_math/ConditioningMathNode.py +++ b/more_math/ConditioningMathNode.py @@ -1,6 +1,6 @@ from unittest import result import torch -from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like +from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like, get_v_variable from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io from antlr4 import InputStream, CommonTokenStream @@ -135,6 +135,12 @@ class ConditioningMathNode(io.ComfyNode): "batch_count": a.shape[0], } | generate_dim_variables(a) | V_norm_tensors + v_stacked, v_cnt = get_v_variable(V_norm_tensors, length_mismatch=length_mismatch) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + for k, val in F.items(): variables[k] = val if val is not None else 0.0 @@ -164,6 +170,12 @@ class ConditioningMathNode(io.ComfyNode): "batch_count": a_p.shape[0] if a_p.numel() > 0 else 0, } | generate_dim_variables(a_p) | V_norm_pooled + v_stacked, v_cnt = get_v_variable(V_norm_pooled, length_mismatch=length_mismatch) + if v_stacked is not None: + variables_pi["V"] = v_stacked + variables_pi["Vcnt"] = float(v_cnt) + variables_pi["V_count"] = float(v_cnt) + for k, val in F.items(): variables_pi[k] = val if val is not None else 0.0 diff --git a/more_math/FloatMathNode.py b/more_math/FloatMathNode.py index 903a9ba..9a8b15e 100644 --- a/more_math/FloatMathNode.py +++ b/more_math/FloatMathNode.py @@ -1,7 +1,7 @@ from inspect import cleandoc import torch -from .helper_functions import parse_expr +from .helper_functions import parse_expr, get_v_variable from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io @@ -86,6 +86,12 @@ class FloatMathNode(io.ComfyNode): for k, val in V.items(): variables[k] = val if val is not None else 0.0 + v_stacked, v_cnt = get_v_variable(variables) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + tree = parse_expr(FloatFunc); # scalar execution # UnifiedMathVisitor expects variables and a shape. Shape [1] for scalar? diff --git a/more_math/GuiderMathNode.py b/more_math/GuiderMathNode.py index 2885a65..7d5c3aa 100644 --- a/more_math/GuiderMathNode.py +++ b/more_math/GuiderMathNode.py @@ -10,6 +10,7 @@ from .helper_functions import ( parse_expr, make_zero_like, as_tensor, + get_v_variable ) from comfy_api.latest import io import comfy.sampler_helpers @@ -165,6 +166,12 @@ class MathGuider: "c": g_results.get("V2", make_zero_like(eval_samples)), "d": g_results.get("V3", make_zero_like(eval_samples)), }) + + v_stacked, v_cnt = get_v_variable(g_results) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) for k, v in self.F.items(): variables[k] = v if v is not None else 0.0 diff --git a/more_math/ImageMathNode.py b/more_math/ImageMathNode.py index 5b1bca4..006eeb2 100644 --- a/more_math/ImageMathNode.py +++ b/more_math/ImageMathNode.py @@ -1,4 +1,4 @@ -from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like +from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like, get_v_variable from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io from antlr4 import InputStream, CommonTokenStream @@ -128,6 +128,12 @@ class ImageMathNode(io.ComfyNode): # Add all dynamic inputs variables.update(V_norm) + v_stacked, v_cnt = get_v_variable(V_norm, length_mismatch=length_mismatch) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + for k, val in F.items(): variables[k] = val if val is not None else 0.0 diff --git a/more_math/LatentMathNode.py b/more_math/LatentMathNode.py index 03638a4..7356fc3 100644 --- a/more_math/LatentMathNode.py +++ b/more_math/LatentMathNode.py @@ -6,7 +6,8 @@ from .helper_functions import ( parse_expr, as_tensor, normalize_to_common_shape, - make_zero_like + make_zero_like, + get_v_variable ) from .Parser.UnifiedMathVisitor import UnifiedMathVisitor import torch @@ -181,6 +182,12 @@ class LatentMathNode(io.ComfyNode): # Add all dynamic inputs variables.update(V_norm_samples) + v_stacked, v_cnt = get_v_variable(V_norm_samples, length_mismatch=length_mismatch) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + for k, v in F.items(): variables[k] = v if v is not None else 0.0 diff --git a/more_math/MaskMathNode.py b/more_math/MaskMathNode.py index 73acf85..bcfed7d 100644 --- a/more_math/MaskMathNode.py +++ b/more_math/MaskMathNode.py @@ -1,4 +1,4 @@ -from .helper_functions import generate_dim_variables,parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape,make_zero_like +from .helper_functions import generate_dim_variables,parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape,make_zero_like, get_v_variable from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io from antlr4 import InputStream, CommonTokenStream @@ -120,6 +120,12 @@ class MaskMathNode(io.ComfyNode): "batch_count": ae.shape[0], } | generate_dim_variables(ae) + v_stacked, v_cnt = get_v_variable(V_norm, length_mismatch=length_mismatch) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + # Add all dynamic inputs variables.update(V_norm) diff --git a/more_math/NoiseMathNode.py b/more_math/NoiseMathNode.py index 545b9c5..1438cdb 100644 --- a/more_math/NoiseMathNode.py +++ b/more_math/NoiseMathNode.py @@ -1,4 +1,4 @@ -from .helper_functions import generate_dim_variables, as_tensor, parse_expr, getIndexTensorAlongDim, make_zero_like +from .helper_functions import generate_dim_variables, as_tensor, parse_expr, getIndexTensorAlongDim, make_zero_like, get_v_variable from comfy_api.latest import io import torch from .Parser.MathExprParser import MathExprParser,InputStream,CommonTokenStream @@ -123,6 +123,12 @@ class NoiseExecutor: "input_latent": samples, } | generate_dim_variables(samples) | vals | self.F + v_stacked, v_cnt = get_v_variable(vals) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + if time_dim is not None: F = getIndexTensorAlongDim(samples, time_dim) variables.update({"frame": F, "frame_count": frame_count}) diff --git a/more_math/VideoMathNode.py b/more_math/VideoMathNode.py index 864e7d5..880dfae 100644 --- a/more_math/VideoMathNode.py +++ b/more_math/VideoMathNode.py @@ -1,5 +1,5 @@ import torch -from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like +from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like, get_v_variable from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from comfy_api.latest import io from antlr4 import InputStream, CommonTokenStream @@ -130,6 +130,12 @@ class VideoMathNode(io.ComfyNode): "channel_count": ae.shape[3], } | generate_dim_variables(ae) + v_stacked, v_cnt = get_v_variable(V_norm, length_mismatch=length_mismatch) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + # Add all dynamic inputs variables.update(V_norm) @@ -187,6 +193,13 @@ class VideoMathNode(io.ComfyNode): "batch_count": a_w.shape[0], } | generate_dim_variables(a_w) | V_norm_waveforms | sample_rates + v_stacked, v_cnt = get_v_variable(V_norm_waveforms, length_mismatch=length_mismatch) + if v_stacked is not None: + # This 'variables' is the one for the second eval in VideoMathNode + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + for k, val in F.items(): variables[k] = val if val is not None else 0.0 diff --git a/more_math/helper_functions.py b/more_math/helper_functions.py index 39ae649..9cfe82e 100644 --- a/more_math/helper_functions.py +++ b/more_math/helper_functions.py @@ -194,3 +194,36 @@ def normalize_to_common_shape(*tensors, mode="pad"): else: result.append(t.contiguous()) return tuple(result) + +def get_v_variable(v_norm_dict, length_mismatch="error"): + """ + Collects V0, V1, ... from the dict, stacks them into a V tensor, + and returns (V_stacked, V_count). + """ + sorted_keys = sorted([k for k in v_norm_dict.keys() if k.startswith("V")], key=lambda x: int(x[1:])) + ordered_tensors = [] + + for k in sorted_keys: + val = v_norm_dict[k] + if torch.is_tensor(val): + ordered_tensors.append(val) + elif isinstance(val, (int, float)): + ordered_tensors.append(torch.tensor(val)) + + if not ordered_tensors: + return None, 0 + + if length_mismatch == "error" and len(ordered_tensors) > 1: + first_shape = ordered_tensors[0].shape + for i in range(1, len(ordered_tensors)): + if ordered_tensors[i].shape != first_shape: + raise ValueError(f"Input variables have mismatched shapes: {first_shape} vs {ordered_tensors[i].shape}. Cannot create 'V' variable. Switch 'length_mismatch' to 'pad' or 'tile' to enable 'V'.") + + try: + stacked = torch.stack(ordered_tensors) + return stacked, len(ordered_tensors) + except RuntimeError as e: + if length_mismatch == "error": + raise ValueError(f"Failed to stack input variables into 'V': {str(e)}") + return None, len(ordered_tensors) + diff --git a/more_math/modelLikeCommon.py b/more_math/modelLikeCommon.py index 143397e..ff89ade 100644 --- a/more_math/modelLikeCommon.py +++ b/more_math/modelLikeCommon.py @@ -1,4 +1,4 @@ -from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor +from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, get_v_variable from .Parser.UnifiedMathVisitor import UnifiedMathVisitor import torch @@ -102,6 +102,12 @@ def calculate_patches_autogrow(Expr, V, F, mapping=None): elif target in V: # V exists but key missing variables[alias] = torch.zeros_like(ref_tensor) + v_stacked, v_cnt = get_v_variable(variables) + if v_stacked is not None: + variables["V"] = v_stacked + variables["Vcnt"] = float(v_cnt) + variables["V_count"] = float(v_cnt) + variables = variables | generate_dim_variables(ref_tensor) # Execute math diff --git a/repro_grammar.py b/repro_grammar.py new file mode 100644 index 0000000..63632d6 --- /dev/null +++ b/repro_grammar.py @@ -0,0 +1,35 @@ +from antlr4 import InputStream, CommonTokenStream +from more_math.Parser.MathExprLexer import MathExprLexer +from more_math.Parser.MathExprParser import MathExprParser +import sys + +def test_expr(expr): + print(f"Testing: {expr}") + input_stream = InputStream(expr) + lexer = MathExprLexer(input_stream) + # Redirect stderr to catch ANTLR errors + stream = CommonTokenStream(lexer) + parser = MathExprParser(stream) + + class ErrorCountListener: + def __init__(self): + self.count = 0 + def syntaxError(self, recognizer, offendingSymbol, line, column, msg, e): + print(f"Error at {line}:{column}: {msg}") + self.count += 1 + + listener = ErrorCountListener() + parser.removeErrorListeners() + parser.addErrorListener(listener) + + tree = parser.start() + if listener.count > 0: + print("FAILED") + else: + print("PASSED") + +if __name__ == "__main__": + test_expr("1+1") + test_expr("1+1;") + test_expr("x = 1; x") + test_expr("x = 1; x;")