AI fix stack and large refactor + add count
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
+539
-529
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+1337
-1221
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user