add tensor creation and fix ruff

This commit is contained in:
mcDandy
2026-02-03 10:17:43 +01:00
parent 9d0f87e43f
commit e01a74d6ad
19 changed files with 1996 additions and 1909 deletions
-2
View File
@@ -1,4 +1,3 @@
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, get_v_variable, get_f_variable
from .Parser.UnifiedMathVisitor import UnifiedMathVisitor
@@ -105,7 +104,6 @@ class ConditioningMathNode(io.ComfyNode):
V_norm_tensors = dict(zip(tensor_keys, norm_tensors_batch))
ref_tensor = norm_tensors_batch[0]
common_shape = ref_tensor.shape
# Normalize pooled outputs (if they exist)
valid_pooled_keys = [k for k, v in pooled_outputs.items() if v is not None]
-1
View File
@@ -1,4 +1,3 @@
from numpy import stack
import torch
import re
from antlr4 import InputStream, CommonTokenStream
+5 -2
View File
@@ -1,7 +1,7 @@
grammar MathExpr;
// Top-level entry point
start: (funcDef | varDef | stmt)* expr? SEMICOLON? EOF;
start: (funcDef | varDef | stmt)* expr SEMICOLON? EOF;
funcDef:
VARIABLE LPAREN paramList? RPAREN ARROW (block | expr) SEMICOLON # FunctionDef;
@@ -170,7 +170,8 @@ func2:
| BOTK_IND LPAREN expr COMMA expr RPAREN # BotkIndFunc
| BOTK_IND LPAREN expr COMMA expr RPAREN # BotkIndFunc
| PUSH LPAREN expr COMMA expr RPAREN # PushFunc
| GET_VALUE LPAREN expr COMMA expr RPAREN # GetValueFunc;
| GET_VALUE LPAREN expr COMMA expr RPAREN # GetValueFunc
| TENSOR LPAREN indexExpr (COMMA expr)? RPAREN # EmptyTensorFunc;
func3:
CLAMP LPAREN expr COMMA expr COMMA expr RPAREN # ClampFunc
@@ -322,6 +323,8 @@ NONE: 'None' | 'none' | 'NULL' | 'null';
BREAK: 'break';
CONTINUE: 'continue';
TENSOR: 'tensor';
PLUS: '+';
MINUS: '-';
MULT: '*';
File diff suppressed because one or more lines are too long
+58 -56
View File
@@ -108,37 +108,38 @@ TIMESTAMP=107
NONE=108
BREAK=109
CONTINUE=110
PLUS=111
MINUS=112
MULT=113
DIV=114
MOD=115
POW=116
GE=117
GT=118
LE=119
LT=120
EQ=121
EQUEALS=122
NE=123
PIPE=124
LPAREN=125
RPAREN=126
COMMA=127
SEMICOLON=128
ARROW=129
LBRACKET=130
RBRACKET=131
QUESTION=132
COLON=133
LBRACE=134
RBRACE=135
NUMBER=136
CONSTANT=137
VARIABLE=138
SL_COMMENT=139
ML_COMMENT=140
WS=141
TENSOR=111
PLUS=112
MINUS=113
MULT=114
DIV=115
MOD=116
POW=117
GE=118
GT=119
LE=120
LT=121
EQ=122
EQUEALS=123
NE=124
PIPE=125
LPAREN=126
RPAREN=127
COMMA=128
SEMICOLON=129
ARROW=130
LBRACKET=131
RBRACKET=132
QUESTION=133
COLON=134
LBRACE=135
RBRACE=136
NUMBER=137
CONSTANT=138
VARIABLE=139
SL_COMMENT=140
ML_COMMENT=141
WS=142
'sin'=1
'cos'=2
'tan'=3
@@ -224,28 +225,29 @@ WS=141
'in'=106
'break'=109
'continue'=110
'+'=111
'-'=112
'*'=113
'/'=114
'%'=115
'^'=116
'>='=117
'>'=118
'<='=119
'<'=120
'=='=121
'='=122
'!='=123
'|'=124
'('=125
')'=126
','=127
';'=128
'->'=129
'['=130
']'=131
'?'=132
':'=133
'{'=134
'}'=135
'tensor'=111
'+'=112
'-'=113
'*'=114
'/'=115
'%'=116
'^'=117
'>='=118
'>'=119
'<='=120
'<'=121
'=='=122
'='=123
'!='=124
'|'=125
'('=126
')'=127
','=128
';'=129
'->'=130
'['=131
']'=132
'?'=133
':'=134
'{'=135
'}'=136
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+58 -56
View File
@@ -108,37 +108,38 @@ TIMESTAMP=107
NONE=108
BREAK=109
CONTINUE=110
PLUS=111
MINUS=112
MULT=113
DIV=114
MOD=115
POW=116
GE=117
GT=118
LE=119
LT=120
EQ=121
EQUEALS=122
NE=123
PIPE=124
LPAREN=125
RPAREN=126
COMMA=127
SEMICOLON=128
ARROW=129
LBRACKET=130
RBRACKET=131
QUESTION=132
COLON=133
LBRACE=134
RBRACE=135
NUMBER=136
CONSTANT=137
VARIABLE=138
SL_COMMENT=139
ML_COMMENT=140
WS=141
TENSOR=111
PLUS=112
MINUS=113
MULT=114
DIV=115
MOD=116
POW=117
GE=118
GT=119
LE=120
LT=121
EQ=122
EQUEALS=123
NE=124
PIPE=125
LPAREN=126
RPAREN=127
COMMA=128
SEMICOLON=129
ARROW=130
LBRACKET=131
RBRACKET=132
QUESTION=133
COLON=134
LBRACE=135
RBRACE=136
NUMBER=137
CONSTANT=138
VARIABLE=139
SL_COMMENT=140
ML_COMMENT=141
WS=142
'sin'=1
'cos'=2
'tan'=3
@@ -224,28 +225,29 @@ WS=141
'in'=106
'break'=109
'continue'=110
'+'=111
'-'=112
'*'=113
'/'=114
'%'=115
'^'=116
'>='=117
'>'=118
'<='=119
'<'=120
'=='=121
'='=122
'!='=123
'|'=124
'('=125
')'=126
','=127
';'=128
'->'=129
'['=130
']'=131
'?'=132
':'=133
'{'=134
'}'=135
'tensor'=111
'+'=112
'-'=113
'*'=114
'/'=115
'%'=116
'^'=117
'>='=118
'>'=119
'<='=120
'<'=121
'=='=122
'='=123
'!='=124
'|'=125
'('=126
')'=127
','=128
';'=129
'->'=130
'['=131
']'=132
'?'=133
':'=134
'{'=135
'}'=136
+9
View File
@@ -1259,6 +1259,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#EmptyTensorFunc.
def enterEmptyTensorFunc(self, ctx:MathExprParser.EmptyTensorFuncContext):
pass
# Exit a parse tree produced by MathExprParser#EmptyTensorFunc.
def exitEmptyTensorFunc(self, ctx:MathExprParser.EmptyTensorFuncContext):
pass
# Enter a parse tree produced by MathExprParser#ClampFunc.
def enterClampFunc(self, ctx:MathExprParser.ClampFuncContext):
pass
File diff suppressed because it is too large Load Diff
+5
View File
@@ -704,6 +704,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#EmptyTensorFunc.
def visitEmptyTensorFunc(self, ctx:MathExprParser.EmptyTensorFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ClampFunc.
def visitClampFunc(self, ctx:MathExprParser.ClampFuncContext):
return self.visitChildren(ctx)
+10 -3
View File
@@ -50,7 +50,7 @@ class UnifiedMathVisitor(MathExprVisitor):
if isinstance(last_result, (ReturnSignal, BreakSignal, ContinueSignal)):
parent_gen = stack[-1]
func_name = parent_gen.gi_code.co_name
is_handler = False
if isinstance(last_result, (BreakSignal, ContinueSignal)):
if func_name in ("visitWhileStmt", "visitForStmt"):
@@ -58,7 +58,7 @@ class UnifiedMathVisitor(MathExprVisitor):
elif isinstance(last_result, ReturnSignal):
if func_name in ("visitCallExp", "visitStart"):
is_handler = True
if not is_handler:
stack.pop().close()
continue
@@ -1839,6 +1839,8 @@ class UnifiedMathVisitor(MathExprVisitor):
return res
def visitEdgeFunc(self, ctx):
tsr_val = yield ctx.expr(0)
tsr = self._promote_to_tensor(tsr_val)
original_shape = tsr.shape
tsr = tsr.float()
@@ -1962,4 +1964,9 @@ class UnifiedMathVisitor(MathExprVisitor):
return BreakSignal()
def visitContinueExp(self, ctx):
return ContinueSignal()
return ContinueSignal()
def visitEmptyTensorFunc(self, ctx):
value = (yield ctx.expr()) if ctx.expr() else 0.0
shape = yield ctx.indexExpr()
return torch.full(shape, value, device=self.device)
-1
View File
@@ -3,7 +3,6 @@ from comfy_api.latest import io
import copy
from antlr4 import InputStream, CommonTokenStream
from custom_nodes.more_math.more_math.Stack import MrmthStack
from .Parser.MathExprLexer import MathExprLexer
from .Parser.MathExprParser import MathExprParser
import re
-1
View File
@@ -1,4 +1,3 @@
import torch
from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, normalize_to_common_shape, make_zero_like, get_v_variable, get_f_variable
from .Parser.UnifiedMathVisitor import UnifiedMathVisitor
from comfy_api.latest import io
+5 -5
View File
@@ -202,14 +202,14 @@ def get_v_variable(v_norm_dict, length_mismatch="error"):
"""
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
@@ -234,7 +234,7 @@ def get_f_variable(f_dict):
"""
sorted_keys = sorted([k for k in f_dict.keys() if k.startswith("F")], key=lambda x: int(x[1:]))
ordered_values = []
for k in sorted_keys:
val = f_dict[k]
if torch.is_tensor(val):
@@ -245,10 +245,10 @@ def get_f_variable(f_dict):
ordered_values.append(torch.tensor(float(val)))
else:
ordered_values.append(torch.tensor(0.0))
if not ordered_values:
return None, 0
try:
stacked = torch.stack(ordered_values)
return stacked, len(ordered_values)
+1 -2
View File
@@ -1,8 +1,7 @@
from .helper_functions import generate_dim_variables, parse_expr, getIndexTensorAlongDim, as_tensor, get_v_variable, get_f_variable
from .helper_functions import generate_dim_variables, parse_expr, as_tensor, get_v_variable, get_f_variable
from .Parser.UnifiedMathVisitor import UnifiedMathVisitor
import torch
from custom_nodes.more_math.more_math import helper_functions
def calculate_patches(Model, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0):
"""Legacy calculate_patches for backward compatibility."""
-1
View File
@@ -1,4 +1,3 @@
import torch
import sys
import os
+12 -12
View File
@@ -4,14 +4,14 @@ from .test_unified_math import parse_and_visit
def test_multidimensional_indexing():
T = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
vars = {"T": T}
# 1. Single index (dim 0)
assert torch.allclose(parse_and_visit("T[0]", vars), torch.tensor([1.0, 2.0]))
# 2. Multi index
assert parse_and_visit("T[1, 0]", vars) == 3.0
assert parse_and_visit("T[0, 1]", vars) == 2.0
# 3. List indexing on dim 0
res_list = parse_and_visit("T[[0, 1]]", vars)
assert torch.allclose(res_list, T)
@@ -19,15 +19,15 @@ def test_multidimensional_indexing():
def test_tensor_assignment():
T = torch.zeros((2, 2))
vars = {"T": T}
# 1. Scalar to position
parse_and_visit("T[0, 0] = 5.0;", vars)
assert T[0, 0] == 5.0
# 2. Scalar to slice (broadcast)
parse_and_visit("T[1] = 3.0;", vars)
assert torch.allclose(T[1], torch.tensor([3.0, 3.0]))
# 3. Tensor to slice
parse_and_visit("T[0] = [1, 2];", vars)
assert torch.allclose(T[0], torch.tensor([1.0, 2.0]))
@@ -36,12 +36,12 @@ def test_4d_assignment():
# val[0] expect 4d if dim0=1 or 3d
T = torch.zeros((2, 3, 4, 4))
vars = {"T": T}
# Slice assignment
val_3d = torch.ones((3, 4, 4))
parse_and_visit("T[0] = V1;", {"T": T, "V1": val_3d})
assert torch.allclose(T[0], val_3d)
# 4D with leading 1 assignment
val_4d_1 = torch.ones((1, 3, 4, 4)) * 2.0
parse_and_visit("T[1] = V2;", {"T": T, "V2": val_4d_1})
@@ -54,11 +54,11 @@ def test_4d_assignment():
def test_list_assignment():
L = [[1, 2], [3, 4]]
vars = {"L": L}
# Nested assignment
parse_and_visit("L[0, 1] = 99;", vars)
assert L[0][1] == 99
# Multi-bracket syntax
parse_and_visit("L[1][0] = 88;", vars)
assert L[1][0] == 88
@@ -69,11 +69,11 @@ def test_enhanced_assignment_logic():
vars = {"T": T}
parse_and_visit("T[0] = 5.0;", vars)
assert torch.all(T[0] == 5.0)
# 2. 1-element tensor/list filling
parse_and_visit("T[1] = [7.0];", vars)
assert torch.all(T[1] == 7.0)
# 3. Rank matching (squeeze leading 1s)
# Target slice T[0, 0] is 1D (shape [2])
# Value is 3D (shape [1, 1, 2])
+7 -7
View File
@@ -399,7 +399,7 @@ def test_random_generators():
res_p = parse_and_visit("randp(123, 5.0)", vars)
assert res_p.shape == (1, 1, 1, 1)
assert torch.all(res_p >= 0)
def test_recursion_and_depth():
vars = {}
# 1. Test recursion depth (100 levels)
@@ -423,17 +423,17 @@ def test_recursion_and_depth():
def test_new_loop_features():
vars = {}
# 1. Test FOR loop with range() (list)
# x = 0; for(i in range(0, 5, 1)) x = x + i; x
expr1 = "x = 0; for(i in range(0, 5, 1)) x = x + i; x"
assert parse_and_visit(expr1, vars) == 10.0
# 2. Test FOR loop with list
# x = 1; for(i in [1, 2, 3]) x = x * i; x
expr2 = "x = 1; for(i in [1, 2, 3]) x = x * i; x"
assert parse_and_visit(expr2, vars) == 6.0
# 3. Test get_value (2D tensor)
# T = [[1, 2], [3, 4]] -> pos=[1, 0] -> 3
t = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
@@ -442,7 +442,7 @@ def test_new_loop_features():
expr3 = "get_value(T, [1, 0])"
res3 = parse_and_visit(expr3, vars)
assert res3 == 3.0
# get_value(T, [0, 1]) -> 2
assert parse_and_visit("get_value(T, [0, 1])", vars) == 2.0
@@ -450,11 +450,11 @@ def test_new_loop_features():
# crop(T, [0, 0], [1, 1]) -> [[1]]
res4 = parse_and_visit("crop(T, [0, 0], [1, 1])", vars)
assert torch.equal(res4, torch.tensor([[1.0]]))
# crop(T, [0, 0], [2, 2]) -> T
res5 = parse_and_visit("crop(T, [0, 0], [2, 2])", vars)
assert torch.equal(res5, t)
# 5. Test crop with padding (zeros)
# crop(T, [1, 1], [2, 2]) -> [[4, 0], [0, 0]]
# T at [1,1] is 4. Size [2,2].