add shape function

This commit is contained in:
mcDandy
2026-02-15 15:47:52 +01:00
parent 0604551199
commit e1f37d25e3
10 changed files with 2477 additions and 2408 deletions
+1
View File
@@ -152,6 +152,7 @@ You can also get the node from comfy manager under the name of More math.
- `cumprod(x)`: Returns the cumulative product of elements along the batch dimension (dim 0).
- `tensor(shape,value)`: Createss a tensor of given shape filled with value. Value can be omittend and defaults to zero.
- `flatten(value)`: Flattens a tensor to 1D. If input is list, it flattens nested lists into a single list.
- `shape(value)` : Returns the shape of a tensor as a tensor. If input is a list, returns lenght of the list as 1 value tensor. For numbers it returns empty tensor.
### Advanced Tensor Operations
+3 -1
View File
@@ -156,7 +156,8 @@ func1:
| MOTION_MASK LPAREN expr RPAREN # MotionMaskFunc
| FLOW_TO_IMAGE LPAREN expr RPAREN # FlowToImageFunc
| BNOT LPAREN expr RPAREN # BitNotFunc
| BITCOUNT LPAREN expr RPAREN # BitCountFunc;
| BITCOUNT LPAREN expr RPAREN # BitCountFunc
| SHAPE LPAREN expr RPAREN # ShapeFunc;
func2:
POWE LPAREN expr COMMA expr RPAREN # PowFunc
@@ -335,6 +336,7 @@ MATMUL: 'matmul';
RIFE: 'rife';
BNOT: 'bnot' | 'bitwise_not';
BITCOUNT: 'bitcount' | 'popcount' | 'popcnt';
SHAPE: 'shape';
BAND: 'band' | 'bitwise_and';
XOR: 'bxor' | 'bitwise_xor';
BOR: 'bor' | 'bitwise_or';
File diff suppressed because one or more lines are too long
+131 -129
View File
@@ -93,82 +93,83 @@ MATMUL=92
RIFE=93
BNOT=94
BITCOUNT=95
BAND=96
XOR=97
BOR=98
TENSOR=99
PUSH=100
POP=101
CLEAR=102
HAS=103
GET=104
IF=105
ELSE=106
WHILE=107
FOR=108
IN=109
BREAK=110
CONTINUE=111
RETURN=112
TIMESTAMP=113
SORT=114
ARGSORT=115
ARGMIN=116
ARGMAX=117
SOFTMAX=118
SOFTMIN=119
UNIQUE=120
FLIP=121
COV=122
CROP=123
NONE=124
NOISE=125
RAND=126
CAUCHY=127
EXPONENTIAL=128
LOGNORMAL=129
BERNOULLI=130
POISSON=131
GAMMADIST=132
BETADIST=133
LAPLACEDIST=134
GUMBELDIST=135
WEIBULLDIST=136
CHI2DIST=137
STUDENTTDIST=138
PLUS=139
MINUS=140
MULT=141
DIV=142
MOD=143
POW=144
LSHIFT=145
RSHIFT=146
GE=147
GT=148
LE=149
LT=150
EQ=151
EQUEALS=152
NE=153
PIPE=154
LPAREN=155
RPAREN=156
COMMA=157
SEMICOLON=158
ARROW=159
LBRACKET=160
RBRACKET=161
QUESTION=162
COLON=163
LBRACE=164
RBRACE=165
NUMBER=166
CONSTANT=167
VARIABLE=168
SL_COMMENT=169
ML_COMMENT=170
WS=171
SHAPE=96
BAND=97
XOR=98
BOR=99
TENSOR=100
PUSH=101
POP=102
CLEAR=103
HAS=104
GET=105
IF=106
ELSE=107
WHILE=108
FOR=109
IN=110
BREAK=111
CONTINUE=112
RETURN=113
TIMESTAMP=114
SORT=115
ARGSORT=116
ARGMIN=117
ARGMAX=118
SOFTMAX=119
SOFTMIN=120
UNIQUE=121
FLIP=122
COV=123
CROP=124
NONE=125
NOISE=126
RAND=127
CAUCHY=128
EXPONENTIAL=129
LOGNORMAL=130
BERNOULLI=131
POISSON=132
GAMMADIST=133
BETADIST=134
LAPLACEDIST=135
GUMBELDIST=136
WEIBULLDIST=137
CHI2DIST=138
STUDENTTDIST=139
PLUS=140
MINUS=141
MULT=142
DIV=143
MOD=144
POW=145
LSHIFT=146
RSHIFT=147
GE=148
GT=149
LE=150
LT=151
EQ=152
EQUEALS=153
NE=154
PIPE=155
LPAREN=156
RPAREN=157
COMMA=158
SEMICOLON=159
ARROW=160
LBRACKET=161
RBRACKET=162
QUESTION=163
COLON=164
LBRACE=165
RBRACE=166
NUMBER=167
CONSTANT=168
VARIABLE=169
SL_COMMENT=170
ML_COMMENT=171
WS=172
'sin'=1
'cos'=2
'tan'=3
@@ -244,56 +245,57 @@ WS=171
'cross'=91
'matmul'=92
'rife'=93
'tensor'=99
'stack_push'=100
'stack_pop'=101
'stack_clear'=102
'stack_has'=103
'stack_get'=104
'if'=105
'else'=106
'while'=107
'for'=108
'in'=109
'break'=110
'continue'=111
'return'=112
'timestamp'=113
'sort'=114
'argsort'=115
'argmin'=116
'argmax'=117
'softmax'=118
'softmin'=119
'unique'=120
'flip'=121
'cov'=122
'crop'=123
'none'=124
'+'=139
'-'=140
'*'=141
'/'=142
'%'=143
'^'=144
'<<'=145
'>>'=146
'>='=147
'>'=148
'<='=149
'<'=150
'=='=151
'='=152
'!='=153
'|'=154
'('=155
')'=156
','=157
';'=158
'->'=159
'['=160
']'=161
'?'=162
':'=163
'{'=164
'}'=165
'shape'=96
'tensor'=100
'stack_push'=101
'stack_pop'=102
'stack_clear'=103
'stack_has'=104
'stack_get'=105
'if'=106
'else'=107
'while'=108
'for'=109
'in'=110
'break'=111
'continue'=112
'return'=113
'timestamp'=114
'sort'=115
'argsort'=116
'argmin'=117
'argmax'=118
'softmax'=119
'softmin'=120
'unique'=121
'flip'=122
'cov'=123
'crop'=124
'none'=125
'+'=140
'-'=141
'*'=142
'/'=143
'%'=144
'^'=145
'<<'=146
'>>'=147
'>='=148
'>'=149
'<='=150
'<'=151
'=='=152
'='=153
'!='=154
'|'=155
'('=156
')'=157
','=158
';'=159
'->'=160
'['=161
']'=162
'?'=163
':'=164
'{'=165
'}'=166
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+131 -129
View File
@@ -93,82 +93,83 @@ MATMUL=92
RIFE=93
BNOT=94
BITCOUNT=95
BAND=96
XOR=97
BOR=98
TENSOR=99
PUSH=100
POP=101
CLEAR=102
HAS=103
GET=104
IF=105
ELSE=106
WHILE=107
FOR=108
IN=109
BREAK=110
CONTINUE=111
RETURN=112
TIMESTAMP=113
SORT=114
ARGSORT=115
ARGMIN=116
ARGMAX=117
SOFTMAX=118
SOFTMIN=119
UNIQUE=120
FLIP=121
COV=122
CROP=123
NONE=124
NOISE=125
RAND=126
CAUCHY=127
EXPONENTIAL=128
LOGNORMAL=129
BERNOULLI=130
POISSON=131
GAMMADIST=132
BETADIST=133
LAPLACEDIST=134
GUMBELDIST=135
WEIBULLDIST=136
CHI2DIST=137
STUDENTTDIST=138
PLUS=139
MINUS=140
MULT=141
DIV=142
MOD=143
POW=144
LSHIFT=145
RSHIFT=146
GE=147
GT=148
LE=149
LT=150
EQ=151
EQUEALS=152
NE=153
PIPE=154
LPAREN=155
RPAREN=156
COMMA=157
SEMICOLON=158
ARROW=159
LBRACKET=160
RBRACKET=161
QUESTION=162
COLON=163
LBRACE=164
RBRACE=165
NUMBER=166
CONSTANT=167
VARIABLE=168
SL_COMMENT=169
ML_COMMENT=170
WS=171
SHAPE=96
BAND=97
XOR=98
BOR=99
TENSOR=100
PUSH=101
POP=102
CLEAR=103
HAS=104
GET=105
IF=106
ELSE=107
WHILE=108
FOR=109
IN=110
BREAK=111
CONTINUE=112
RETURN=113
TIMESTAMP=114
SORT=115
ARGSORT=116
ARGMIN=117
ARGMAX=118
SOFTMAX=119
SOFTMIN=120
UNIQUE=121
FLIP=122
COV=123
CROP=124
NONE=125
NOISE=126
RAND=127
CAUCHY=128
EXPONENTIAL=129
LOGNORMAL=130
BERNOULLI=131
POISSON=132
GAMMADIST=133
BETADIST=134
LAPLACEDIST=135
GUMBELDIST=136
WEIBULLDIST=137
CHI2DIST=138
STUDENTTDIST=139
PLUS=140
MINUS=141
MULT=142
DIV=143
MOD=144
POW=145
LSHIFT=146
RSHIFT=147
GE=148
GT=149
LE=150
LT=151
EQ=152
EQUEALS=153
NE=154
PIPE=155
LPAREN=156
RPAREN=157
COMMA=158
SEMICOLON=159
ARROW=160
LBRACKET=161
RBRACKET=162
QUESTION=163
COLON=164
LBRACE=165
RBRACE=166
NUMBER=167
CONSTANT=168
VARIABLE=169
SL_COMMENT=170
ML_COMMENT=171
WS=172
'sin'=1
'cos'=2
'tan'=3
@@ -244,56 +245,57 @@ WS=171
'cross'=91
'matmul'=92
'rife'=93
'tensor'=99
'stack_push'=100
'stack_pop'=101
'stack_clear'=102
'stack_has'=103
'stack_get'=104
'if'=105
'else'=106
'while'=107
'for'=108
'in'=109
'break'=110
'continue'=111
'return'=112
'timestamp'=113
'sort'=114
'argsort'=115
'argmin'=116
'argmax'=117
'softmax'=118
'softmin'=119
'unique'=120
'flip'=121
'cov'=122
'crop'=123
'none'=124
'+'=139
'-'=140
'*'=141
'/'=142
'%'=143
'^'=144
'<<'=145
'>>'=146
'>='=147
'>'=148
'<='=149
'<'=150
'=='=151
'='=152
'!='=153
'|'=154
'('=155
')'=156
','=157
';'=158
'->'=159
'['=160
']'=161
'?'=162
':'=163
'{'=164
'}'=165
'shape'=96
'tensor'=100
'stack_push'=101
'stack_pop'=102
'stack_clear'=103
'stack_has'=104
'stack_get'=105
'if'=106
'else'=107
'while'=108
'for'=109
'in'=110
'break'=111
'continue'=112
'return'=113
'timestamp'=114
'sort'=115
'argsort'=116
'argmin'=117
'argmax'=118
'softmax'=119
'softmin'=120
'unique'=121
'flip'=122
'cov'=123
'crop'=124
'none'=125
'+'=140
'-'=141
'*'=142
'/'=143
'%'=144
'^'=145
'<<'=146
'>>'=147
'>='=148
'>'=149
'<='=150
'<'=151
'=='=152
'='=153
'!='=154
'|'=155
'('=156
')'=157
','=158
';'=159
'->'=160
'['=161
']'=162
'?'=163
':'=164
'{'=165
'}'=166
File diff suppressed because it is too large Load Diff
+5
View File
@@ -644,6 +644,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ShapeFunc.
def visitShapeFunc(self, ctx:MathExprParser.ShapeFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#PowFunc.
def visitPowFunc(self, ctx:MathExprParser.PowFuncContext):
return self.visitChildren(ctx)
+19 -10
View File
@@ -234,7 +234,7 @@ class UnifiedMathVisitor(MathExprVisitor):
indices.append(torch.tensor(idx_val, dtype=torch.long, device=self.device))
else:
indices.append(int(idx_val))
# Use standard PyTorch/list indexing
if self._is_tensor(val):
idx_tuple = tuple(indices)
@@ -252,7 +252,7 @@ class UnifiedMathVisitor(MathExprVisitor):
idx = int(idx.item())
current = current[idx]
return current
error_prefix = f"{ctx.start.line}:{ctx.start.column}:"
raise ValueError(f"{error_prefix} Indexing only supported on tensors and lists (found {type(val).__name__})")
@@ -1878,7 +1878,7 @@ class UnifiedMathVisitor(MathExprVisitor):
shape_arg = self.shape
if len(ctx.expr()) > 3:
shape_arg = (yield ctx.expr(3))
# Use torch.distributions.Gamma which internally handles the generator properly via torch.manual_seed
# We set the random state temporarily
old_state = torch.get_rng_state()
@@ -1900,7 +1900,7 @@ class UnifiedMathVisitor(MathExprVisitor):
shape_arg = self.shape
if len(ctx.expr()) > 3:
shape_arg = (yield ctx.expr(3))
old_state = torch.get_rng_state()
try:
torch.manual_seed(seed)
@@ -1949,7 +1949,7 @@ class UnifiedMathVisitor(MathExprVisitor):
shape_arg = self.shape
if len(ctx.expr()) > 3:
shape_arg = (yield ctx.expr(3))
# Implement Weibull using generator-aware uniform: scale * (-log(u))^(1/concentration)
generator = torch.Generator(device=self.device).manual_seed(seed)
u = torch.rand(shape_arg, generator=generator, device=self.device)
@@ -1963,7 +1963,7 @@ class UnifiedMathVisitor(MathExprVisitor):
shape_arg = self.shape
if len(ctx.expr()) > 2:
shape_arg = (yield ctx.expr(2))
# Chi-squared is Gamma(df/2, 2)
old_state = torch.get_rng_state()
try:
@@ -1982,11 +1982,11 @@ class UnifiedMathVisitor(MathExprVisitor):
shape_arg = self.shape
if len(ctx.expr()) > 2:
shape_arg = (yield ctx.expr(2))
# Student's t using normal and chi-squared: Z / sqrt(V/df) where Z~N(0,1) and V~Chi2(df)
generator = torch.Generator(device=self.device).manual_seed(seed)
z = torch.randn(shape_arg, generator=generator, device=self.device)
# Generate chi-squared using the same seed + 1 to maintain determinism but different samples
old_state = torch.get_rng_state()
try:
@@ -1996,7 +1996,7 @@ class UnifiedMathVisitor(MathExprVisitor):
v = v.to(device=self.device)
finally:
torch.set_rng_state(old_state)
return z / torch.sqrt(v / df)
def visitNvlFunc(self, ctx):
@@ -2372,6 +2372,15 @@ class UnifiedMathVisitor(MathExprVisitor):
v = (yield ctx.expr())
return self._bitwise_popcount(v)
def visitShapeFunc(self, ctx):
val = (yield ctx.expr())
if self._is_tensor(val):
return val.shape.to(self.device)
elif self._is_list(val):
return torch.tensor([len(val)], dtype=torch.long, device=self.device)
else:
return torch.tensor([], dtype=torch.long, device=self.device)
def _bitwise_op(self, a, b, torch_op, scalar_op):
"""Binary bitwise operation handler supporting tensors, lists, and scalars."""
if self._is_tensor(a) and a.numel() == 1:
@@ -2496,7 +2505,7 @@ class UnifiedMathVisitor(MathExprVisitor):
if self._is_tensor(v):
v_t = self._promote_to_tensor(v).flatten().long()
# Use numpy's bin and count for efficiency
counts = torch.tensor([bin(int(x) & 0xFFFFFFFFFFFFFFFF).count('1') for x in v_t.tolist()],
counts = torch.tensor([bin(int(x) & 0xFFFFFFFFFFFFFFFF).count('1') for x in v_t.tolist()],
dtype=torch.float32, device=v_t.device)
if counts.numel() == 1:
return float(counts.item())