argsort + add do nothing option on tensor size mismatch
This commit is contained in:
@@ -38,7 +38,8 @@ class AudioMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on input audio"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
|
||||
@@ -24,7 +24,8 @@ class CLIPMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on weights"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched layer counts. For models, this usually defaults to broadcast (zero for missing layers)."
|
||||
),
|
||||
|
||||
@@ -32,7 +32,8 @@ class ConditioningMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression_pi",display_name="pooled output expr.", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on pooled_input part of conditioning"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
|
||||
@@ -29,7 +29,8 @@ class ImageMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on input images"), # Changed ID to Expression to match AudioMathNode pattern, or keep Image? AudioMathNode used "Expression".
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
|
||||
@@ -40,7 +40,8 @@ class LatentMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on input latents"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched latent batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
|
||||
@@ -30,7 +30,8 @@ class MaskMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on input masks"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched mask batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
|
||||
@@ -24,7 +24,8 @@ class ModelMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on weights"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched layer counts. For models, this usually defaults to broadcast (zero for missing layers)."
|
||||
),
|
||||
|
||||
@@ -143,7 +143,8 @@ func1:
|
||||
| POP LPAREN expr RPAREN # PopFunc
|
||||
| CLEAR LPAREN expr RPAREN # ClearFunc
|
||||
| HAS LPAREN expr RPAREN # HasFunc
|
||||
| GET LPAREN expr RPAREN # GetFunc;
|
||||
| GET LPAREN expr RPAREN # GetFunc
|
||||
| ARGSORT LPAREN expr (COMMA expr)? RPAREN # ArgsortFunc;
|
||||
|
||||
// Two-argument functions Two-argument functions
|
||||
func2:
|
||||
@@ -316,6 +317,7 @@ APPEND: 'append';
|
||||
GET_VALUE: 'get_value';
|
||||
BATCH_SHUFFLE: 'batch_shuffle' | 'shuffle' | 'select';
|
||||
CROP: 'crop';
|
||||
ARGSORT: 'argsort';
|
||||
FOR: 'for';
|
||||
IN: 'in';
|
||||
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -103,44 +103,45 @@ APPEND=102
|
||||
GET_VALUE=103
|
||||
BATCH_SHUFFLE=104
|
||||
CROP=105
|
||||
FOR=106
|
||||
IN=107
|
||||
TIMESTAMP=108
|
||||
NONE=109
|
||||
BREAK=110
|
||||
CONTINUE=111
|
||||
TENSOR=112
|
||||
PLUS=113
|
||||
MINUS=114
|
||||
MULT=115
|
||||
DIV=116
|
||||
MOD=117
|
||||
POW=118
|
||||
GE=119
|
||||
GT=120
|
||||
LE=121
|
||||
LT=122
|
||||
EQ=123
|
||||
EQUEALS=124
|
||||
NE=125
|
||||
PIPE=126
|
||||
LPAREN=127
|
||||
RPAREN=128
|
||||
COMMA=129
|
||||
SEMICOLON=130
|
||||
ARROW=131
|
||||
LBRACKET=132
|
||||
RBRACKET=133
|
||||
QUESTION=134
|
||||
COLON=135
|
||||
LBRACE=136
|
||||
RBRACE=137
|
||||
NUMBER=138
|
||||
CONSTANT=139
|
||||
VARIABLE=140
|
||||
SL_COMMENT=141
|
||||
ML_COMMENT=142
|
||||
WS=143
|
||||
ARGSORT=106
|
||||
FOR=107
|
||||
IN=108
|
||||
TIMESTAMP=109
|
||||
NONE=110
|
||||
BREAK=111
|
||||
CONTINUE=112
|
||||
TENSOR=113
|
||||
PLUS=114
|
||||
MINUS=115
|
||||
MULT=116
|
||||
DIV=117
|
||||
MOD=118
|
||||
POW=119
|
||||
GE=120
|
||||
GT=121
|
||||
LE=122
|
||||
LT=123
|
||||
EQ=124
|
||||
EQUEALS=125
|
||||
NE=126
|
||||
PIPE=127
|
||||
LPAREN=128
|
||||
RPAREN=129
|
||||
COMMA=130
|
||||
SEMICOLON=131
|
||||
ARROW=132
|
||||
LBRACKET=133
|
||||
RBRACKET=134
|
||||
QUESTION=135
|
||||
COLON=136
|
||||
LBRACE=137
|
||||
RBRACE=138
|
||||
NUMBER=139
|
||||
CONSTANT=140
|
||||
VARIABLE=141
|
||||
SL_COMMENT=142
|
||||
ML_COMMENT=143
|
||||
WS=144
|
||||
'sin'=1
|
||||
'cos'=2
|
||||
'tan'=3
|
||||
@@ -222,33 +223,34 @@ WS=143
|
||||
'append'=102
|
||||
'get_value'=103
|
||||
'crop'=105
|
||||
'for'=106
|
||||
'in'=107
|
||||
'break'=110
|
||||
'continue'=111
|
||||
'tensor'=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
|
||||
'}'=137
|
||||
'argsort'=106
|
||||
'for'=107
|
||||
'in'=108
|
||||
'break'=111
|
||||
'continue'=112
|
||||
'tensor'=113
|
||||
'+'=114
|
||||
'-'=115
|
||||
'*'=116
|
||||
'/'=117
|
||||
'%'=118
|
||||
'^'=119
|
||||
'>='=120
|
||||
'>'=121
|
||||
'<='=122
|
||||
'<'=123
|
||||
'=='=124
|
||||
'='=125
|
||||
'!='=126
|
||||
'|'=127
|
||||
'('=128
|
||||
')'=129
|
||||
','=130
|
||||
';'=131
|
||||
'->'=132
|
||||
'['=133
|
||||
']'=134
|
||||
'?'=135
|
||||
':'=136
|
||||
'{'=137
|
||||
'}'=138
|
||||
|
||||
File diff suppressed because one or more lines are too long
+565
-559
File diff suppressed because it is too large
Load Diff
@@ -103,44 +103,45 @@ APPEND=102
|
||||
GET_VALUE=103
|
||||
BATCH_SHUFFLE=104
|
||||
CROP=105
|
||||
FOR=106
|
||||
IN=107
|
||||
TIMESTAMP=108
|
||||
NONE=109
|
||||
BREAK=110
|
||||
CONTINUE=111
|
||||
TENSOR=112
|
||||
PLUS=113
|
||||
MINUS=114
|
||||
MULT=115
|
||||
DIV=116
|
||||
MOD=117
|
||||
POW=118
|
||||
GE=119
|
||||
GT=120
|
||||
LE=121
|
||||
LT=122
|
||||
EQ=123
|
||||
EQUEALS=124
|
||||
NE=125
|
||||
PIPE=126
|
||||
LPAREN=127
|
||||
RPAREN=128
|
||||
COMMA=129
|
||||
SEMICOLON=130
|
||||
ARROW=131
|
||||
LBRACKET=132
|
||||
RBRACKET=133
|
||||
QUESTION=134
|
||||
COLON=135
|
||||
LBRACE=136
|
||||
RBRACE=137
|
||||
NUMBER=138
|
||||
CONSTANT=139
|
||||
VARIABLE=140
|
||||
SL_COMMENT=141
|
||||
ML_COMMENT=142
|
||||
WS=143
|
||||
ARGSORT=106
|
||||
FOR=107
|
||||
IN=108
|
||||
TIMESTAMP=109
|
||||
NONE=110
|
||||
BREAK=111
|
||||
CONTINUE=112
|
||||
TENSOR=113
|
||||
PLUS=114
|
||||
MINUS=115
|
||||
MULT=116
|
||||
DIV=117
|
||||
MOD=118
|
||||
POW=119
|
||||
GE=120
|
||||
GT=121
|
||||
LE=122
|
||||
LT=123
|
||||
EQ=124
|
||||
EQUEALS=125
|
||||
NE=126
|
||||
PIPE=127
|
||||
LPAREN=128
|
||||
RPAREN=129
|
||||
COMMA=130
|
||||
SEMICOLON=131
|
||||
ARROW=132
|
||||
LBRACKET=133
|
||||
RBRACKET=134
|
||||
QUESTION=135
|
||||
COLON=136
|
||||
LBRACE=137
|
||||
RBRACE=138
|
||||
NUMBER=139
|
||||
CONSTANT=140
|
||||
VARIABLE=141
|
||||
SL_COMMENT=142
|
||||
ML_COMMENT=143
|
||||
WS=144
|
||||
'sin'=1
|
||||
'cos'=2
|
||||
'tan'=3
|
||||
@@ -222,33 +223,34 @@ WS=143
|
||||
'append'=102
|
||||
'get_value'=103
|
||||
'crop'=105
|
||||
'for'=106
|
||||
'in'=107
|
||||
'break'=110
|
||||
'continue'=111
|
||||
'tensor'=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
|
||||
'}'=137
|
||||
'argsort'=106
|
||||
'for'=107
|
||||
'in'=108
|
||||
'break'=111
|
||||
'continue'=112
|
||||
'tensor'=113
|
||||
'+'=114
|
||||
'-'=115
|
||||
'*'=116
|
||||
'/'=117
|
||||
'%'=118
|
||||
'^'=119
|
||||
'>='=120
|
||||
'>'=121
|
||||
'<='=122
|
||||
'<'=123
|
||||
'=='=124
|
||||
'='=125
|
||||
'!='=126
|
||||
'|'=127
|
||||
'('=128
|
||||
')'=129
|
||||
','=130
|
||||
';'=131
|
||||
'->'=132
|
||||
'['=133
|
||||
']'=134
|
||||
'?'=135
|
||||
':'=136
|
||||
'{'=137
|
||||
'}'=138
|
||||
|
||||
@@ -1052,6 +1052,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#ArgsortFunc.
|
||||
def enterArgsortFunc(self, ctx:MathExprParser.ArgsortFuncContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#ArgsortFunc.
|
||||
def exitArgsortFunc(self, ctx:MathExprParser.ArgsortFuncContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#PowFunc.
|
||||
def enterPowFunc(self, ctx:MathExprParser.PowFuncContext):
|
||||
pass
|
||||
|
||||
+805
-742
File diff suppressed because it is too large
Load Diff
@@ -589,6 +589,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#ArgsortFunc.
|
||||
def visitArgsortFunc(self, ctx:MathExprParser.ArgsortFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#PowFunc.
|
||||
def visitPowFunc(self, ctx:MathExprParser.PowFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
@@ -679,6 +679,13 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
|
||||
return tsr[indices]
|
||||
|
||||
def visitArgsortFunc(self, ctx):
|
||||
val = self._promote_to_tensor((yield ctx.expr(0)))
|
||||
descending = False
|
||||
if ctx.expr(1):
|
||||
descending = bool((yield ctx.expr(1)))
|
||||
return torch.argsort(val, descending=descending)
|
||||
|
||||
# Three-argument functions
|
||||
def visitClampFunc(self, ctx):
|
||||
val = (yield ctx.expr(0))
|
||||
@@ -1460,7 +1467,6 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
a = self._promote_to_tensor(a)
|
||||
b = self._promote_to_tensor(b)
|
||||
|
||||
# Ensure at least 1D for cat
|
||||
if a.ndim == 0:
|
||||
a = a.unsqueeze(0)
|
||||
if b.ndim == 0:
|
||||
|
||||
@@ -29,7 +29,8 @@ class SigmasMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on input images"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
|
||||
@@ -25,7 +25,8 @@ class VAEMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on weights"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched layer counts. For models, this usually defaults to broadcast (zero for missing layers)."
|
||||
),
|
||||
|
||||
@@ -30,7 +30,8 @@ class VideoMathNode(io.ComfyNode):
|
||||
io.String.Input(id="Expression_pi", default="I0*(1-F0)+I1*F0", tooltip="Expression to apply on pooled_input part of conditioning"),
|
||||
io.Combo.Input(
|
||||
id="length_mismatch",
|
||||
options=["tile", "error", "pad"],
|
||||
options=["do nothing","error","tile", "pad"],
|
||||
display_name="on size mismatch",
|
||||
default="error",
|
||||
tooltip="How to handle mismatched image batch sizes. tile: repeat shorter inputs; error: raise error on mismatch; pad: treat missing frames as zero."
|
||||
),
|
||||
|
||||
Reference in New Issue
Block a user