AI fix stack and large refactor + add count

This commit is contained in:
mcDandy
2026-02-02 20:16:22 +01:00
parent a1d1ce8c03
commit d53d725fcc
14 changed files with 2139 additions and 1930 deletions
+1
View File
@@ -178,6 +178,7 @@ You can also get the node from comfy manager under the name of More math.
- `nan_to_num(x, nan_value, posinf_value, neginf_value)` or `nvl`: Replaces NaN and infinite values in tensor with specified values.
- `remap(v, i_min, i_max, o_min, o_max)`: Remaps value `v` from input range `[i_min, i_max]` to output range `[o_min, o_max]`.
- `timestamp()` or `now`: Returns current UNIX timestamp (precision to microseconds, can be different on other systems)
- `count(x)` or `length(x)` or `cnt(x)`: Returns the length of a list or the size of the first dimension of a tensor.
### Random Distributions
+1 -1
View File
@@ -89,7 +89,7 @@ class LatentMathNode(io.ComfyNode):
return needed1
@classmethod
def execute(cls, V, F, Expression,batching, length_mismatch="tile",stack=[]) -> io.NodeOutput:
def execute(cls, V, F, Expression,batching, length_mismatch="tile",stack=dict()) -> io.NodeOutput:
# Determine reference latent
ref_latent = None
for lat in V.values():
+1 -1
View File
@@ -72,7 +72,7 @@ class ModelMathNode(io.ComfyNode):
return needed1
@classmethod
def execute(cls, V, F, Expression, length_mismatch="tile",stack=[]) -> io.NodeOutput:
def execute(cls, V, F, Expression, length_mismatch="tile",stack=dict()) -> io.NodeOutput:
# Determine reference model for cloning
a = V.get("V0")
if a is None:
+6 -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;
@@ -81,7 +81,9 @@ atom:
| PIPE expr PIPE # AbsExp
| LBRACKET expr (COMMA expr)* RBRACKET # ListExp
| VARIABLE LPAREN exprList? RPAREN # CallExp
| NONE # NoneExp;
| NONE # NoneExp
| BREAK # BreakExp
| CONTINUE # ContinueExp;
exprList: expr (COMMA expr)*;
@@ -136,6 +138,7 @@ func1:
| MEDIAN LPAREN expr RPAREN # MedianFunc
| MODE LPAREN expr RPAREN # ModeFunc
| CUMSUM LPAREN expr RPAREN # CumsumFunc
| COUNT LPAREN expr RPAREN # CountFunc
| CUMPROD LPAREN expr RPAREN # CumprodFunc
| POP LPAREN expr RPAREN # PopFunc
| CLEAR LPAREN expr RPAREN # ClearFunc
@@ -307,6 +310,7 @@ COSSIM: 'cossim';
FLIP: 'flip';
COV: 'cov';
SORT: 'sort';
COUNT: 'count' | 'length' | 'cnt';
APPEND: 'append';
GET_VALUE: 'get_value';
CROP: 'crop';
File diff suppressed because one or more lines are too long
+73 -72
View File
@@ -98,46 +98,47 @@ COSSIM=97
FLIP=98
COV=99
SORT=100
APPEND=101
GET_VALUE=102
CROP=103
FOR=104
IN=105
TIMESTAMP=106
NONE=107
BREAK=108
CONTINUE=109
PLUS=110
MINUS=111
MULT=112
DIV=113
MOD=114
POW=115
GE=116
GT=117
LE=118
LT=119
EQ=120
EQUEALS=121
NE=122
PIPE=123
LPAREN=124
RPAREN=125
COMMA=126
SEMICOLON=127
ARROW=128
LBRACKET=129
RBRACKET=130
QUESTION=131
COLON=132
LBRACE=133
RBRACE=134
NUMBER=135
CONSTANT=136
VARIABLE=137
SL_COMMENT=138
ML_COMMENT=139
WS=140
COUNT=101
APPEND=102
GET_VALUE=103
CROP=104
FOR=105
IN=106
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
'sin'=1
'cos'=2
'tan'=3
@@ -216,35 +217,35 @@ WS=140
'flip'=98
'cov'=99
'sort'=100
'append'=101
'get_value'=102
'crop'=103
'for'=104
'in'=105
'break'=108
'continue'=109
'+'=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
'append'=102
'get_value'=103
'crop'=104
'for'=105
'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
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+73 -72
View File
@@ -98,46 +98,47 @@ COSSIM=97
FLIP=98
COV=99
SORT=100
APPEND=101
GET_VALUE=102
CROP=103
FOR=104
IN=105
TIMESTAMP=106
NONE=107
BREAK=108
CONTINUE=109
PLUS=110
MINUS=111
MULT=112
DIV=113
MOD=114
POW=115
GE=116
GT=117
LE=118
LT=119
EQ=120
EQUEALS=121
NE=122
PIPE=123
LPAREN=124
RPAREN=125
COMMA=126
SEMICOLON=127
ARROW=128
LBRACKET=129
RBRACKET=130
QUESTION=131
COLON=132
LBRACE=133
RBRACE=134
NUMBER=135
CONSTANT=136
VARIABLE=137
SL_COMMENT=138
ML_COMMENT=139
WS=140
COUNT=101
APPEND=102
GET_VALUE=103
CROP=104
FOR=105
IN=106
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
'sin'=1
'cos'=2
'tan'=3
@@ -216,35 +217,35 @@ WS=140
'flip'=98
'cov'=99
'sort'=100
'append'=101
'get_value'=102
'crop'=103
'for'=104
'in'=105
'break'=108
'continue'=109
'+'=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
'append'=102
'get_value'=103
'crop'=104
'for'=105
'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
+27
View File
@@ -530,6 +530,24 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#BreakExp.
def enterBreakExp(self, ctx:MathExprParser.BreakExpContext):
pass
# Exit a parse tree produced by MathExprParser#BreakExp.
def exitBreakExp(self, ctx:MathExprParser.BreakExpContext):
pass
# Enter a parse tree produced by MathExprParser#ContinueExp.
def enterContinueExp(self, ctx:MathExprParser.ContinueExpContext):
pass
# Exit a parse tree produced by MathExprParser#ContinueExp.
def exitContinueExp(self, ctx:MathExprParser.ContinueExpContext):
pass
# Enter a parse tree produced by MathExprParser#exprList.
def enterExprList(self, ctx:MathExprParser.ExprListContext):
pass
@@ -980,6 +998,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#CountFunc.
def enterCountFunc(self, ctx:MathExprParser.CountFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CountFunc.
def exitCountFunc(self, ctx:MathExprParser.CountFuncContext):
pass
# Enter a parse tree produced by MathExprParser#CumprodFunc.
def enterCumprodFunc(self, ctx:MathExprParser.CumprodFuncContext):
pass
File diff suppressed because it is too large Load Diff
+15
View File
@@ -299,6 +299,16 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#BreakExp.
def visitBreakExp(self, ctx:MathExprParser.BreakExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ContinueExp.
def visitContinueExp(self, ctx:MathExprParser.ContinueExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#exprList.
def visitExprList(self, ctx:MathExprParser.ExprListContext):
return self.visitChildren(ctx)
@@ -549,6 +559,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#CountFunc.
def visitCountFunc(self, ctx:MathExprParser.CountFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#CumprodFunc.
def visitCumprodFunc(self, ctx:MathExprParser.CumprodFuncContext):
return self.visitChildren(ctx)
+58 -29
View File
@@ -46,12 +46,24 @@ class UnifiedMathVisitor(MathExprVisitor):
while stack:
try:
res = stack[-1].send(last_result)
# If we're bubbling a signal, we need to check if the parent can handle it.
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"):
is_handler = True
elif isinstance(last_result, ReturnSignal):
if func_name in ("visitCallExp", "visitStart"):
is_handler = True
if not is_handler:
stack.pop().close()
continue
if isinstance(res, ReturnSignal):
stack.pop()
last_result = res
continue
res = stack[-1].send(last_result)
if hasattr(res, 'accept'):
next_gen = res.accept(self)
@@ -956,6 +968,16 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitSumFunc(self, ctx):
return self._reduction_op((yield ctx.expr()), torch.sum, sum)
def visitCountFunc(self, ctx):
val = yield ctx.expr()
if self._is_list(val):
return float(len(val))
if self._is_tensor(val):
if val.ndim == 0:
return 1.0
return float(val.size(0))
return 1.0
def visitMeanFunc(self, ctx):
return self._reduction_op(
(yield ctx.expr()), lambda x: torch.mean(x.float()), lambda x: sum(x) / len(x) if x else 0.0
@@ -1423,11 +1445,8 @@ class UnifiedMathVisitor(MathExprVisitor):
continue
res = yield child
if isinstance(res, ReturnSignal):
return res.value
if isinstance(res, (BreakSignal, ContinueSignal)):
raise RuntimeError("break/continue outside of loop")
last_res = res
if res is not None:
last_res = res
return last_res
@@ -1446,11 +1465,10 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitBlock(self, ctx):
vars_before = set(self.variables.keys())
try:
val = None
for stmt in ctx.stmt():
res = yield stmt
if isinstance(res, (ReturnSignal, BreakSignal, ContinueSignal)):
return res
return None
val = yield stmt
return val
finally:
for v in set(self.variables.keys()) - vars_before:
del self.variables[v]
@@ -1469,15 +1487,10 @@ class UnifiedMathVisitor(MathExprVisitor):
return bool(x)
if truthy(cond):
res = yield ctx.stmt(0)
return (yield ctx.stmt(0))
elif ctx.stmt(1):
res = yield ctx.stmt(1)
else:
return None
if isinstance(res, ReturnSignal):
return res
return res
return (yield ctx.stmt(1))
return None
def visitWhileStmt(self, ctx):
while True:
@@ -1520,8 +1533,6 @@ class UnifiedMathVisitor(MathExprVisitor):
self.variables[var_name] = val
res = yield ctx.stmt()
if isinstance(res, ReturnSignal):
return res
if isinstance(res, BreakSignal):
break
if isinstance(res, ContinueSignal):
@@ -1903,7 +1914,15 @@ class UnifiedMathVisitor(MathExprVisitor):
return o_min + (v - i_min) * (o_max - o_min) / denom
def _ensure_dict_storage(self):
if not isinstance(self._state_storage, dict):
if not self._state_storage:
self._state_storage = {}
else:
self._state_storage = {i: v for i, v in enumerate(self._state_storage)}
def visitPushFunc(self, ctx):
self._ensure_dict_storage()
f= yield ctx.expr(0)
slot = int(f)
if slot not in self._state_storage:
@@ -1913,24 +1932,34 @@ class UnifiedMathVisitor(MathExprVisitor):
return value
def visitPopFunc(self, ctx):
slot = yield ctx.expr()
self._ensure_dict_storage()
slot = int((yield ctx.expr()))
if slot not in self._state_storage or not self._state_storage[slot]:
raise ValueError(f"Pop from empty slot: {slot}")
return self._state_storage[slot].pop()
def visitClearFunc(self, ctx):
slot = yield ctx.expr()
self._ensure_dict_storage()
slot = int((yield ctx.expr()))
if slot in self._state_storage:
self._state_storage[slot] = []
return None
def visitHasFunc(self, ctx):
slot = yield ctx.expr()
self._ensure_dict_storage()
slot = int((yield ctx.expr()))
return float(slot in self._state_storage and bool(self._state_storage[slot]))
def visitGetFunc(self, ctx):
slot = yield ctx.expr()
self._ensure_dict_storage()
slot = int((yield ctx.expr()))
if slot not in self._state_storage:
raise ValueError(f"Get from empty slot: {slot}")
storage_list = self._state_storage[slot]
return storage_list.last()
return storage_list[-1] if storage_list else None
def visitBreakExp(self, ctx):
return BreakSignal()
def visitContinueExp(self, ctx):
return ContinueSignal()
+1 -1
View File
@@ -73,7 +73,7 @@ class VAEMathNode(io.ComfyNode):
return needed1
@classmethod
def execute(cls, V, F, Expression, length_mismatch="tile",stack=[]) -> io.NodeOutput:
def execute(cls, V, F, Expression, length_mismatch="tile",stack=dict()) -> io.NodeOutput:
# Determine reference VAE
a = V.get("V0")
if a is None: