add tensor creation and fix ruff
This commit is contained in:
@@ -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,4 +1,3 @@
|
||||
from numpy import stack
|
||||
import torch
|
||||
import re
|
||||
from antlr4 import InputStream, CommonTokenStream
|
||||
|
||||
@@ -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
@@ -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
+498
-494
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+1321
-1264
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -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,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,4 +1,3 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
|
||||
|
||||
+12
-12
@@ -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])
|
||||
|
||||
@@ -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].
|
||||
|
||||
Reference in New Issue
Block a user