argsort + add do nothing option on tensor size mismatch

This commit is contained in:
mcDandy
2026-02-04 19:09:56 +01:00
parent 8ffc9a394c
commit 0c1d3f52d6
20 changed files with 1561 additions and 1451 deletions
+2 -1
View File
@@ -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."
),
+2 -1
View File
@@ -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)."
),
+2 -1
View File
@@ -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."
),
+2 -1
View File
@@ -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."
),
+2 -1
View File
@@ -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."
),
+2 -1
View File
@@ -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."
),
+2 -1
View File
@@ -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)."
),
+3 -1
View File
@@ -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
+70 -68
View File
@@ -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
File diff suppressed because it is too large Load Diff
+70 -68
View File
@@ -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
+9
View File
@@ -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
File diff suppressed because it is too large Load Diff
+5
View File
@@ -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)
+7 -1
View File
@@ -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:
+2 -1
View File
@@ -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."
),
+2 -1
View File
@@ -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)."
),
+2 -1
View File
@@ -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."
),