add shape function
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
+768
-764
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
+1412
-1373
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
|
||||
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user