add cord func. Same as Dnumber but for any tensor

This commit is contained in:
mcDandy
2026-08-26 13:29:22 +02:00
parent b783d6f770
commit 736275ef28
22 changed files with 5284 additions and 5052 deletions
+1
View File
@@ -367,6 +367,7 @@ ___
### Generators
- `text_image(text, font, size, [max_width], [weight], [rotation_angle], [line_spacing], [italic], [underline])` - renders text to an 2D tensor
- `coordinates(shape, dim, [dtype])` / `coords` - generates a tensor whose values are the coordinates of each element along the specified dimension. dtype is copied from tensor at that position. Default is float32.
---
+6 -1
View File
@@ -707,7 +707,11 @@ func3:
/**
text_image(text, font, size, [max_width], [weight], [rotation_angle], [line_spacing], [italic], [underline]) - renders text to an 2D tensor
*/
| TEXT_IMAGE LPAREN expr COMMA expr COMMA expr (COMMA expr)? (COMMA expr)? (COMMA expr)? (COMMA expr)? (COMMA expr)? (COMMA expr)? RPAREN # TextImageFunc;
| TEXT_IMAGE LPAREN expr COMMA expr COMMA expr (COMMA expr)? (COMMA expr)? (COMMA expr)? (COMMA expr)? (COMMA expr)? (COMMA expr)? RPAREN # TextImageFunc
/**
coords(shape, dim, [dtype]) - generates a tensor with the specified shape whose values are the coordinates of each element along the specified dimension. The dtype can be specified to control the data type of the output tensor.
*/
| COORDS LPAREN expr COMMA expr (COMMA expr)? RPAREN # CoordsFunc;
func4:
/**
@@ -1035,6 +1039,7 @@ CORR: 'corr' | 'correlation';
ENTROPY: 'entropy';
CROP: 'crop';
NONE: 'none'|'None'|'null'|'NULL';
COORDS: 'coords'|'coordinates';
NOISE: 'noise' | 'randn' | 'random_normal';
RAND: 'rand' | 'randu' | 'random_uniform';
+8 -1
View File
@@ -4973,4 +4973,11 @@ class UnifiedMathVisitor(MathExprVisitor):
diag_view = torch.diagonal(res, offset=offset, dim1=dim1, dim2=dim2)
diag_view.copy_(x)
return res
return res
def visitCoordsFunc(self, ctx):
"""coords(shape) - generate a grid of coordinates"""
shape = self._get_shape_from_ctx(ctx, 0)
dim = ctx.expr(1)
dtype = ctx.expr(2).dtype if len(ctx.expr()) > 2 else torch.float32
return getIndexTensorAlongDim(torch.zeros(shape, dtype=dtype, device=self.device), dim)
File diff suppressed because one or more lines are too long
+91 -90
View File
@@ -170,64 +170,65 @@ CORR=169
ENTROPY=170
CROP=171
NONE=172
NOISE=173
RAND=174
CAUCHY=175
EXPONENTIAL=176
LOGNORMAL=177
BERNOULLI=178
POISSON=179
GAMMADIST=180
BETADIST=181
LAPLACEDIST=182
GUMBELDIST=183
WEIBULLDIST=184
CHI2DIST=185
STUDENTTDIST=186
PERLIN=187
CELLULAR=188
PLASMA=189
RIDGED=190
DOMAIN_WARP=191
PLUS=192
MINUS=193
MULT=194
DIV=195
MOD=196
POW=197
LSHIFT=198
RSHIFT=199
GE=200
GT=201
LE=202
LT=203
EQ=204
EQUEALS=205
PLUS_EQ=206
MINUS_EQ=207
MULT_EQ=208
DIV_EQ=209
MOD_EQ=210
NE=211
PIPE=212
LPAREN=213
RPAREN=214
COMMA=215
SEMICOLON=216
ARROW=217
LBRACKET=218
RBRACKET=219
QUESTION=220
COLON=221
LBRACE=222
RBRACE=223
NUMBER=224
CONSTANT=225
STRING=226
VARIABLE=227
SL_COMMENT=228
ML_COMMENT=229
WS=230
COORDS=173
NOISE=174
RAND=175
CAUCHY=176
EXPONENTIAL=177
LOGNORMAL=178
BERNOULLI=179
POISSON=180
GAMMADIST=181
BETADIST=182
LAPLACEDIST=183
GUMBELDIST=184
WEIBULLDIST=185
CHI2DIST=186
STUDENTTDIST=187
PERLIN=188
CELLULAR=189
PLASMA=190
RIDGED=191
DOMAIN_WARP=192
PLUS=193
MINUS=194
MULT=195
DIV=196
MOD=197
POW=198
LSHIFT=199
RSHIFT=200
GE=201
GT=202
LE=203
LT=204
EQ=205
EQUEALS=206
PLUS_EQ=207
MINUS_EQ=208
MULT_EQ=209
DIV_EQ=210
MOD_EQ=211
NE=212
PIPE=213
LPAREN=214
RPAREN=215
COMMA=216
SEMICOLON=217
ARROW=218
LBRACKET=219
RBRACKET=220
QUESTION=221
COLON=222
LBRACE=223
RBRACE=224
NUMBER=225
CONSTANT=226
STRING=227
VARIABLE=228
SL_COMMENT=229
ML_COMMENT=230
WS=231
'sin'=1
'cos'=2
'tan'=3
@@ -366,35 +367,35 @@ WS=230
'cov'=168
'entropy'=170
'crop'=171
'+'=192
'-'=193
'*'=194
'/'=195
'%'=196
'^'=197
'<<'=198
'>>'=199
'>='=200
'>'=201
'<='=202
'<'=203
'=='=204
'='=205
'+='=206
'-='=207
'*='=208
'/='=209
'%='=210
'!='=211
'|'=212
'('=213
')'=214
','=215
';'=216
'->'=217
'['=218
']'=219
'?'=220
':'=221
'{'=222
'}'=223
'+'=193
'-'=194
'*'=195
'/'=196
'%'=197
'^'=198
'<<'=199
'>>'=200
'>='=201
'>'=202
'<='=203
'<'=204
'=='=205
'='=206
'+='=207
'-='=208
'*='=209
'/='=210
'%='=211
'!='=212
'|'=213
'('=214
')'=215
','=216
';'=217
'->'=218
'['=219
']'=220
'?'=221
':'=222
'{'=223
'}'=224
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+91 -90
View File
@@ -170,64 +170,65 @@ CORR=169
ENTROPY=170
CROP=171
NONE=172
NOISE=173
RAND=174
CAUCHY=175
EXPONENTIAL=176
LOGNORMAL=177
BERNOULLI=178
POISSON=179
GAMMADIST=180
BETADIST=181
LAPLACEDIST=182
GUMBELDIST=183
WEIBULLDIST=184
CHI2DIST=185
STUDENTTDIST=186
PERLIN=187
CELLULAR=188
PLASMA=189
RIDGED=190
DOMAIN_WARP=191
PLUS=192
MINUS=193
MULT=194
DIV=195
MOD=196
POW=197
LSHIFT=198
RSHIFT=199
GE=200
GT=201
LE=202
LT=203
EQ=204
EQUEALS=205
PLUS_EQ=206
MINUS_EQ=207
MULT_EQ=208
DIV_EQ=209
MOD_EQ=210
NE=211
PIPE=212
LPAREN=213
RPAREN=214
COMMA=215
SEMICOLON=216
ARROW=217
LBRACKET=218
RBRACKET=219
QUESTION=220
COLON=221
LBRACE=222
RBRACE=223
NUMBER=224
CONSTANT=225
STRING=226
VARIABLE=227
SL_COMMENT=228
ML_COMMENT=229
WS=230
COORDS=173
NOISE=174
RAND=175
CAUCHY=176
EXPONENTIAL=177
LOGNORMAL=178
BERNOULLI=179
POISSON=180
GAMMADIST=181
BETADIST=182
LAPLACEDIST=183
GUMBELDIST=184
WEIBULLDIST=185
CHI2DIST=186
STUDENTTDIST=187
PERLIN=188
CELLULAR=189
PLASMA=190
RIDGED=191
DOMAIN_WARP=192
PLUS=193
MINUS=194
MULT=195
DIV=196
MOD=197
POW=198
LSHIFT=199
RSHIFT=200
GE=201
GT=202
LE=203
LT=204
EQ=205
EQUEALS=206
PLUS_EQ=207
MINUS_EQ=208
MULT_EQ=209
DIV_EQ=210
MOD_EQ=211
NE=212
PIPE=213
LPAREN=214
RPAREN=215
COMMA=216
SEMICOLON=217
ARROW=218
LBRACKET=219
RBRACKET=220
QUESTION=221
COLON=222
LBRACE=223
RBRACE=224
NUMBER=225
CONSTANT=226
STRING=227
VARIABLE=228
SL_COMMENT=229
ML_COMMENT=230
WS=231
'sin'=1
'cos'=2
'tan'=3
@@ -366,35 +367,35 @@ WS=230
'cov'=168
'entropy'=170
'crop'=171
'+'=192
'-'=193
'*'=194
'/'=195
'%'=196
'^'=197
'<<'=198
'>>'=199
'>='=200
'>'=201
'<='=202
'<'=203
'=='=204
'='=205
'+='=206
'-='=207
'*='=208
'/='=209
'%='=210
'!='=211
'|'=212
'('=213
')'=214
','=215
';'=216
'->'=217
'['=218
']'=219
'?'=220
':'=221
'{'=222
'}'=223
'+'=193
'-'=194
'*'=195
'/'=196
'%'=197
'^'=198
'<<'=199
'>>'=200
'>='=201
'>'=202
'<='=203
'<'=204
'=='=205
'='=206
'+='=207
'-='=208
'*='=209
'/='=210
'%='=211
'!='=212
'|'=213
'('=214
')'=215
','=216
';'=217
'->'=218
'['=219
']'=220
'?'=221
':'=222
'{'=223
'}'=224
@@ -1934,6 +1934,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#CoordsFunc.
def enterCoordsFunc(self, ctx:MathExprParser.CoordsFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CoordsFunc.
def exitCoordsFunc(self, ctx:MathExprParser.CoordsFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SwapFunc.
def enterSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
pass
File diff suppressed because it is too large Load Diff
@@ -1079,6 +1079,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#CoordsFunc.
def visitCoordsFunc(self, ctx:MathExprParser.CoordsFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SwapFunc.
def visitSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
return self.visitChildren(ctx)
+3 -1
View File
@@ -6,7 +6,7 @@ INBUILT_CONSTANTS = {
'e', 'pi'
}
INBUILT_FUNCTIONS = {
'abs', 'acos', 'acosh', 'all', 'angle', 'any', 'append', 'argmax', 'argmin', 'argsort', 'as_nested_tensor', 'asin', 'asinh', 'atan', 'atan2', 'atanh', 'band', 'batch_shuffle', 'bitcount', 'bitwise_and', 'bitwise_not', 'bitwise_or', 'bitwise_xor', 'blur', 'bnot', 'bor', 'botk', 'botk_ind', 'botk_indices', 'bxor', 'cat', 'ceil', 'cellular', 'cellular_noise', 'cielab_to_rgb', 'clamp', 'cnt', 'concat', 'concatenate', 'conv', 'convolution', 'corr', 'correlation', 'cos', 'cosh', 'cosine_similarity', 'cossim', 'count', 'cov', 'crop', 'cross', 'cubic', 'cubic_ease', 'cumprod', 'cumsum', 'diag', 'diagonal_matrix', 'dilate', 'dist', 'distance', 'domain_warp', 'domain_warp_noise', 'dot', 'edge', 'elastic', 'elastic_ease', 'endswith', 'entropy', 'erf', 'erfinv', 'erode', 'exp', 'ezconv', 'ezconvolution', 'fft', 'find', 'flatten', 'flip', 'float', 'floor', 'flow_ang', 'flow_angle', 'flow_apply', 'flow_mag', 'flow_magnitude', 'flow_to_image', 'fract', 'gamma', 'gaussian', 'gelu', 'get_value', 'hist', 'histogram', 'hsv_to_rgb', 'ifft', 'int', 'int_to_rgb', 'interpolate_area', 'interpolate_linear', 'interpolate_nearest', 'interpolate_nearest_exact', 'join', 'length', 'lerp', 'linspace', 'ln', 'log', 'logspace', 'lower', 'map', 'matmul', 'mean', 'median', 'mode', 'moment', 'morph_close', 'morph_open', 'motion_mask', 'nan_to_num', 'noise', 'now', 'nvl', 'oklab_to_rgb', 'overlay', 'pad', 'percentile', 'perlin', 'perlin_noise', 'perm', 'permute', 'pinv', 'plasma', 'plasma_noise', 'popcnt', 'popcount', 'pow', 'prcnt', 'print', 'print_shape', 'pshp', 'quantile', 'quartil', 'quartile', 'rand', 'randb', 'randbeta', 'randc', 'rande', 'randg', 'randgumbel', 'randchi2', 'randl', 'randln', 'randn', 'random_bernoulli', 'random_beta', 'random_cauchy', 'random_exponential', 'random_gamma', 'random_gumbel', 'random_chi2', 'random_laplace', 'random_log_normal', 'random_normal', 'random_poisson', 'random_studentt', 'random_uniform', 'random_weibull', 'randp', 'randt', 'randu', 'randw', 'range', 'relu', 'remap', 'repeat', 'replace', 'reshape', 'rgb_to_cielab', 'rgb_to_hsv', 'rgb_to_int', 'rgb_to_oklab', 'ridged', 'ridged_noise', 'rife', 'roll', 'round', 'rshp', 'select', 'shape', 'shuffle', 'sigm', 'sign', 'sin', 'sine', 'sine_ease', 'singular_value_decomposition', 'sinh', 'smax', 'smin', 'smootherstep', 'smoothstep', 'snorm', 'softmax', 'softmin', 'softplus', 'sort', 'split', 'sqrt', 'squeeze', 'stack_clear', 'stack_get', 'stack_has', 'stack_pop', 'stack_push', 'startswith', 'std', 'step', 'substr', 'substring', 'sum', 'svd', 'swap', 'tan', 'tanh', 'tensor', 'text_image', 'timestamp', 'tmax', 'tmin', 'tnorm', 'topk', 'topk_ind', 'topk_indices', 'trim', 'turbulence', 'unique', 'unsqueeze', 'upper', 'var', 'voronoi', 'voronoi_noise', 'where', 'worley'
'abs', 'acos', 'acosh', 'all', 'angle', 'any', 'append', 'argmax', 'argmin', 'argsort', 'as_nested_tensor', 'asin', 'asinh', 'atan', 'atan2', 'atanh', 'band', 'batch_shuffle', 'bitcount', 'bitwise_and', 'bitwise_not', 'bitwise_or', 'bitwise_xor', 'blur', 'bnot', 'bor', 'botk', 'botk_ind', 'botk_indices', 'bxor', 'cat', 'ceil', 'cellular', 'cellular_noise', 'cielab_to_rgb', 'clamp', 'cnt', 'concat', 'concatenate', 'conv', 'convolution', 'coordinates', 'coords', 'corr', 'correlation', 'cos', 'cosh', 'cosine_similarity', 'cossim', 'count', 'cov', 'crop', 'cross', 'cubic', 'cubic_ease', 'cumprod', 'cumsum', 'diag', 'diagonal_matrix', 'dilate', 'dist', 'distance', 'domain_warp', 'domain_warp_noise', 'dot', 'edge', 'elastic', 'elastic_ease', 'endswith', 'entropy', 'erf', 'erfinv', 'erode', 'exp', 'ezconv', 'ezconvolution', 'fft', 'find', 'flatten', 'flip', 'float', 'floor', 'flow_ang', 'flow_angle', 'flow_apply', 'flow_mag', 'flow_magnitude', 'flow_to_image', 'fract', 'gamma', 'gaussian', 'gelu', 'get_value', 'hist', 'histogram', 'hsv_to_rgb', 'ifft', 'int', 'int_to_rgb', 'interpolate_area', 'interpolate_linear', 'interpolate_nearest', 'interpolate_nearest_exact', 'join', 'length', 'lerp', 'linspace', 'ln', 'log', 'logspace', 'lower', 'map', 'matmul', 'mean', 'median', 'mode', 'moment', 'morph_close', 'morph_open', 'motion_mask', 'nan_to_num', 'noise', 'now', 'nvl', 'oklab_to_rgb', 'overlay', 'pad', 'percentile', 'perlin', 'perlin_noise', 'perm', 'permute', 'pinv', 'plasma', 'plasma_noise', 'popcnt', 'popcount', 'pow', 'prcnt', 'print', 'print_shape', 'pshp', 'quantile', 'quartil', 'quartile', 'rand', 'randb', 'randbeta', 'randc', 'rande', 'randg', 'randgumbel', 'randchi2', 'randl', 'randln', 'randn', 'random_bernoulli', 'random_beta', 'random_cauchy', 'random_exponential', 'random_gamma', 'random_gumbel', 'random_chi2', 'random_laplace', 'random_log_normal', 'random_normal', 'random_poisson', 'random_studentt', 'random_uniform', 'random_weibull', 'randp', 'randt', 'randu', 'randw', 'range', 'relu', 'remap', 'repeat', 'replace', 'reshape', 'rgb_to_cielab', 'rgb_to_hsv', 'rgb_to_int', 'rgb_to_oklab', 'ridged', 'ridged_noise', 'rife', 'roll', 'round', 'rshp', 'select', 'shape', 'shuffle', 'sigm', 'sign', 'sin', 'sine', 'sine_ease', 'singular_value_decomposition', 'sinh', 'smax', 'smin', 'smootherstep', 'smoothstep', 'snorm', 'softmax', 'softmin', 'softplus', 'sort', 'split', 'sqrt', 'squeeze', 'stack_clear', 'stack_get', 'stack_has', 'stack_pop', 'stack_push', 'startswith', 'std', 'step', 'substr', 'substring', 'sum', 'svd', 'swap', 'tan', 'tanh', 'tensor', 'text_image', 'timestamp', 'tmax', 'tmin', 'tnorm', 'topk', 'topk_ind', 'topk_indices', 'trim', 'turbulence', 'unique', 'unsqueeze', 'upper', 'var', 'voronoi', 'voronoi_noise', 'where', 'worley'
}
INBUILT_FUNCTION_META = {
@@ -51,6 +51,8 @@ INBUILT_FUNCTION_META = {
'concatenate': {'min_args': 2, 'max_args': None, 'snippet': 'concatenate()', 'description': 'concatenate(x1, x2, ..., dim) - concatenates tensors along the specified dimension'},
'conv': {'min_args': 2, 'max_args': None, 'snippet': 'conv()', 'description': 'convolution(tensor, [kernel_sizes...], kernel) - applies convolution with kernel. Expects [batch,channel, ...]'},
'convolution': {'min_args': 2, 'max_args': None, 'snippet': 'convolution()', 'description': 'convolution(tensor, [kernel_sizes...], kernel) - applies convolution with kernel. Expects [batch,channel, ...]'},
'coordinates': {'min_args': 2, 'max_args': 3, 'snippet': 'coordinates()', 'description': 'coords(shape, dim, [dtype]) - generates a tensor with the specified shape whose values are the coordinates of each element along the specified dimension. The dtype can be specified to control the data type of the output tensor.'},
'coords': {'min_args': 2, 'max_args': 3, 'snippet': 'coords()', 'description': 'coords(shape, dim, [dtype]) - generates a tensor with the specified shape whose values are the coordinates of each element along the specified dimension. The dtype can be specified to control the data type of the output tensor.'},
'corr': {'min_args': 2, 'max_args': 2, 'snippet': 'corr()', 'description': 'correlation(x, y) - computes the correlation between x and y'},
'correlation': {'min_args': 2, 'max_args': 2, 'snippet': 'correlation()', 'description': 'correlation(x, y) - computes the correlation between x and y'},
'cos': {'min_args': 1, 'max_args': 1, 'snippet': 'cos()', 'description': 'cos(x) - applies cosinus function to value or each element of value'},
File diff suppressed because one or more lines are too long
+91 -90
View File
@@ -170,64 +170,65 @@ CORR=169
ENTROPY=170
CROP=171
NONE=172
NOISE=173
RAND=174
CAUCHY=175
EXPONENTIAL=176
LOGNORMAL=177
BERNOULLI=178
POISSON=179
GAMMADIST=180
BETADIST=181
LAPLACEDIST=182
GUMBELDIST=183
WEIBULLDIST=184
CHI2DIST=185
STUDENTTDIST=186
PERLIN=187
CELLULAR=188
PLASMA=189
RIDGED=190
DOMAIN_WARP=191
PLUS=192
MINUS=193
MULT=194
DIV=195
MOD=196
POW=197
LSHIFT=198
RSHIFT=199
GE=200
GT=201
LE=202
LT=203
EQ=204
EQUEALS=205
PLUS_EQ=206
MINUS_EQ=207
MULT_EQ=208
DIV_EQ=209
MOD_EQ=210
NE=211
PIPE=212
LPAREN=213
RPAREN=214
COMMA=215
SEMICOLON=216
ARROW=217
LBRACKET=218
RBRACKET=219
QUESTION=220
COLON=221
LBRACE=222
RBRACE=223
NUMBER=224
CONSTANT=225
STRING=226
VARIABLE=227
SL_COMMENT=228
ML_COMMENT=229
WS=230
COORDS=173
NOISE=174
RAND=175
CAUCHY=176
EXPONENTIAL=177
LOGNORMAL=178
BERNOULLI=179
POISSON=180
GAMMADIST=181
BETADIST=182
LAPLACEDIST=183
GUMBELDIST=184
WEIBULLDIST=185
CHI2DIST=186
STUDENTTDIST=187
PERLIN=188
CELLULAR=189
PLASMA=190
RIDGED=191
DOMAIN_WARP=192
PLUS=193
MINUS=194
MULT=195
DIV=196
MOD=197
POW=198
LSHIFT=199
RSHIFT=200
GE=201
GT=202
LE=203
LT=204
EQ=205
EQUEALS=206
PLUS_EQ=207
MINUS_EQ=208
MULT_EQ=209
DIV_EQ=210
MOD_EQ=211
NE=212
PIPE=213
LPAREN=214
RPAREN=215
COMMA=216
SEMICOLON=217
ARROW=218
LBRACKET=219
RBRACKET=220
QUESTION=221
COLON=222
LBRACE=223
RBRACE=224
NUMBER=225
CONSTANT=226
STRING=227
VARIABLE=228
SL_COMMENT=229
ML_COMMENT=230
WS=231
'sin'=1
'cos'=2
'tan'=3
@@ -366,35 +367,35 @@ WS=230
'cov'=168
'entropy'=170
'crop'=171
'+'=192
'-'=193
'*'=194
'/'=195
'%'=196
'^'=197
'<<'=198
'>>'=199
'>='=200
'>'=201
'<='=202
'<'=203
'=='=204
'='=205
'+='=206
'-='=207
'*='=208
'/='=209
'%='=210
'!='=211
'|'=212
'('=213
')'=214
','=215
';'=216
'->'=217
'['=218
']'=219
'?'=220
':'=221
'{'=222
'}'=223
'+'=193
'-'=194
'*'=195
'/'=196
'%'=197
'^'=198
'<<'=199
'>>'=200
'>='=201
'>'=202
'<='=203
'<'=204
'=='=205
'='=206
'+='=207
'-='=208
'*='=209
'/='=210
'%='=211
'!='=212
'|'=213
'('=214
')'=215
','=216
';'=217
'->'=218
'['=219
']'=220
'?'=221
':'=222
'{'=223
'}'=224
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+91 -90
View File
@@ -170,64 +170,65 @@ CORR=169
ENTROPY=170
CROP=171
NONE=172
NOISE=173
RAND=174
CAUCHY=175
EXPONENTIAL=176
LOGNORMAL=177
BERNOULLI=178
POISSON=179
GAMMADIST=180
BETADIST=181
LAPLACEDIST=182
GUMBELDIST=183
WEIBULLDIST=184
CHI2DIST=185
STUDENTTDIST=186
PERLIN=187
CELLULAR=188
PLASMA=189
RIDGED=190
DOMAIN_WARP=191
PLUS=192
MINUS=193
MULT=194
DIV=195
MOD=196
POW=197
LSHIFT=198
RSHIFT=199
GE=200
GT=201
LE=202
LT=203
EQ=204
EQUEALS=205
PLUS_EQ=206
MINUS_EQ=207
MULT_EQ=208
DIV_EQ=209
MOD_EQ=210
NE=211
PIPE=212
LPAREN=213
RPAREN=214
COMMA=215
SEMICOLON=216
ARROW=217
LBRACKET=218
RBRACKET=219
QUESTION=220
COLON=221
LBRACE=222
RBRACE=223
NUMBER=224
CONSTANT=225
STRING=226
VARIABLE=227
SL_COMMENT=228
ML_COMMENT=229
WS=230
COORDS=173
NOISE=174
RAND=175
CAUCHY=176
EXPONENTIAL=177
LOGNORMAL=178
BERNOULLI=179
POISSON=180
GAMMADIST=181
BETADIST=182
LAPLACEDIST=183
GUMBELDIST=184
WEIBULLDIST=185
CHI2DIST=186
STUDENTTDIST=187
PERLIN=188
CELLULAR=189
PLASMA=190
RIDGED=191
DOMAIN_WARP=192
PLUS=193
MINUS=194
MULT=195
DIV=196
MOD=197
POW=198
LSHIFT=199
RSHIFT=200
GE=201
GT=202
LE=203
LT=204
EQ=205
EQUEALS=206
PLUS_EQ=207
MINUS_EQ=208
MULT_EQ=209
DIV_EQ=210
MOD_EQ=211
NE=212
PIPE=213
LPAREN=214
RPAREN=215
COMMA=216
SEMICOLON=217
ARROW=218
LBRACKET=219
RBRACKET=220
QUESTION=221
COLON=222
LBRACE=223
RBRACE=224
NUMBER=225
CONSTANT=226
STRING=227
VARIABLE=228
SL_COMMENT=229
ML_COMMENT=230
WS=231
'sin'=1
'cos'=2
'tan'=3
@@ -366,35 +367,35 @@ WS=230
'cov'=168
'entropy'=170
'crop'=171
'+'=192
'-'=193
'*'=194
'/'=195
'%'=196
'^'=197
'<<'=198
'>>'=199
'>='=200
'>'=201
'<='=202
'<'=203
'=='=204
'='=205
'+='=206
'-='=207
'*='=208
'/='=209
'%='=210
'!='=211
'|'=212
'('=213
')'=214
','=215
';'=216
'->'=217
'['=218
']'=219
'?'=220
':'=221
'{'=222
'}'=223
'+'=193
'-'=194
'*'=195
'/'=196
'%'=197
'^'=198
'<<'=199
'>>'=200
'>='=201
'>'=202
'<='=203
'<'=204
'=='=205
'='=206
'+='=207
'-='=208
'*='=209
'/='=210
'%='=211
'!='=212
'|'=213
'('=214
')'=215
','=216
';'=217
'->'=218
'['=219
']'=220
'?'=221
':'=222
'{'=223
'}'=224
@@ -1934,6 +1934,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#CoordsFunc.
def enterCoordsFunc(self, ctx:MathExprParser.CoordsFuncContext):
pass
# Exit a parse tree produced by MathExprParser#CoordsFunc.
def exitCoordsFunc(self, ctx:MathExprParser.CoordsFuncContext):
pass
# Enter a parse tree produced by MathExprParser#SwapFunc.
def enterSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
pass
File diff suppressed because one or more lines are too long
@@ -1079,6 +1079,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#CoordsFunc.
def visitCoordsFunc(self, ctx:MathExprParser.CoordsFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#SwapFunc.
def visitSwapFunc(self, ctx:MathExprParser.SwapFuncContext):
return self.visitChildren(ctx)
+4 -1
View File
@@ -42,7 +42,10 @@ def parse_expr(expr: str):
def getIndexTensorAlongDim(tensor, dim):
"""Create a tensor of indices along a dimension, broadcasted to full shape."""
shape = tensor.shape
if tensor is torch.Tensor:
shape = tensor.shape
else:
shape = torch.Size(tensor)
values = torch.arange(shape[dim], dtype=torch.float32, device=tensor.device)
view_shape = [1] * len(shape)
view_shape[dim] = shape[dim]
+3 -1
View File
@@ -9,7 +9,7 @@ export const CONSTANTS = new Set([
]);
export const FUNCTIONS = new Set([
"abs", "acos", "acosh", "all", "angle", "any", "append", "argmax", "argmin", "argsort", "as_nested_tensor", "asin", "asinh", "atan", "atan2", "atanh", "band", "batch_shuffle", "bitcount", "bitwise_and", "bitwise_not", "bitwise_or", "bitwise_xor", "blur", "bnot", "bor", "botk", "botk_ind", "botk_indices", "bxor", "cat", "ceil", "cellular", "cellular_noise", "cielab_to_rgb", "clamp", "cnt", "concat", "concatenate", "conv", "convolution", "corr", "correlation", "cos", "cosh", "cosine_similarity", "cossim", "count", "cov", "crop", "cross", "cubic", "cubic_ease", "cumprod", "cumsum", "diag", "diagonal_matrix", "dilate", "dist", "distance", "domain_warp", "domain_warp_noise", "dot", "edge", "elastic", "elastic_ease", "endswith", "entropy", "erf", "erfinv", "erode", "exp", "ezconv", "ezconvolution", "fft", "find", "flatten", "flip", "float", "floor", "flow_ang", "flow_angle", "flow_apply", "flow_mag", "flow_magnitude", "flow_to_image", "fract", "gamma", "gaussian", "gelu", "get_value", "hist", "histogram", "hsv_to_rgb", "ifft", "int", "int_to_rgb", "interpolate_area", "interpolate_linear", "interpolate_nearest", "interpolate_nearest_exact", "join", "length", "lerp", "linspace", "ln", "log", "logspace", "lower", "map", "matmul", "mean", "median", "mode", "moment", "morph_close", "morph_open", "motion_mask", "nan_to_num", "noise", "now", "nvl", "oklab_to_rgb", "overlay", "pad", "percentile", "perlin", "perlin_noise", "perm", "permute", "pinv", "plasma", "plasma_noise", "popcnt", "popcount", "pow", "prcnt", "print", "print_shape", "pshp", "quantile", "quartil", "quartile", "rand", "randb", "randbeta", "randc", "rande", "randg", "randgumbel", "randchi2", "randl", "randln", "randn", "random_bernoulli", "random_beta", "random_cauchy", "random_exponential", "random_gamma", "random_gumbel", "random_chi2", "random_laplace", "random_log_normal", "random_normal", "random_poisson", "random_studentt", "random_uniform", "random_weibull", "randp", "randt", "randu", "randw", "range", "relu", "remap", "repeat", "replace", "reshape", "rgb_to_cielab", "rgb_to_hsv", "rgb_to_int", "rgb_to_oklab", "ridged", "ridged_noise", "rife", "roll", "round", "rshp", "select", "shape", "shuffle", "sigm", "sign", "sin", "sine", "sine_ease", "singular_value_decomposition", "sinh", "smax", "smin", "smootherstep", "smoothstep", "snorm", "softmax", "softmin", "softplus", "sort", "split", "sqrt", "squeeze", "stack_clear", "stack_get", "stack_has", "stack_pop", "stack_push", "startswith", "std", "step", "substr", "substring", "sum", "svd", "swap", "tan", "tanh", "tensor", "text_image", "timestamp", "tmax", "tmin", "tnorm", "topk", "topk_ind", "topk_indices", "trim", "turbulence", "unique", "unsqueeze", "upper", "var", "voronoi", "voronoi_noise", "where", "worley"
"abs", "acos", "acosh", "all", "angle", "any", "append", "argmax", "argmin", "argsort", "as_nested_tensor", "asin", "asinh", "atan", "atan2", "atanh", "band", "batch_shuffle", "bitcount", "bitwise_and", "bitwise_not", "bitwise_or", "bitwise_xor", "blur", "bnot", "bor", "botk", "botk_ind", "botk_indices", "bxor", "cat", "ceil", "cellular", "cellular_noise", "cielab_to_rgb", "clamp", "cnt", "concat", "concatenate", "conv", "convolution", "coordinates", "coords", "corr", "correlation", "cos", "cosh", "cosine_similarity", "cossim", "count", "cov", "crop", "cross", "cubic", "cubic_ease", "cumprod", "cumsum", "diag", "diagonal_matrix", "dilate", "dist", "distance", "domain_warp", "domain_warp_noise", "dot", "edge", "elastic", "elastic_ease", "endswith", "entropy", "erf", "erfinv", "erode", "exp", "ezconv", "ezconvolution", "fft", "find", "flatten", "flip", "float", "floor", "flow_ang", "flow_angle", "flow_apply", "flow_mag", "flow_magnitude", "flow_to_image", "fract", "gamma", "gaussian", "gelu", "get_value", "hist", "histogram", "hsv_to_rgb", "ifft", "int", "int_to_rgb", "interpolate_area", "interpolate_linear", "interpolate_nearest", "interpolate_nearest_exact", "join", "length", "lerp", "linspace", "ln", "log", "logspace", "lower", "map", "matmul", "mean", "median", "mode", "moment", "morph_close", "morph_open", "motion_mask", "nan_to_num", "noise", "now", "nvl", "oklab_to_rgb", "overlay", "pad", "percentile", "perlin", "perlin_noise", "perm", "permute", "pinv", "plasma", "plasma_noise", "popcnt", "popcount", "pow", "prcnt", "print", "print_shape", "pshp", "quantile", "quartil", "quartile", "rand", "randb", "randbeta", "randc", "rande", "randg", "randgumbel", "randchi2", "randl", "randln", "randn", "random_bernoulli", "random_beta", "random_cauchy", "random_exponential", "random_gamma", "random_gumbel", "random_chi2", "random_laplace", "random_log_normal", "random_normal", "random_poisson", "random_studentt", "random_uniform", "random_weibull", "randp", "randt", "randu", "randw", "range", "relu", "remap", "repeat", "replace", "reshape", "rgb_to_cielab", "rgb_to_hsv", "rgb_to_int", "rgb_to_oklab", "ridged", "ridged_noise", "rife", "roll", "round", "rshp", "select", "shape", "shuffle", "sigm", "sign", "sin", "sine", "sine_ease", "singular_value_decomposition", "sinh", "smax", "smin", "smootherstep", "smoothstep", "snorm", "softmax", "softmin", "softplus", "sort", "split", "sqrt", "squeeze", "stack_clear", "stack_get", "stack_has", "stack_pop", "stack_push", "startswith", "std", "step", "substr", "substring", "sum", "svd", "swap", "tan", "tanh", "tensor", "text_image", "timestamp", "tmax", "tmin", "tnorm", "topk", "topk_ind", "topk_indices", "trim", "turbulence", "unique", "unsqueeze", "upper", "var", "voronoi", "voronoi_noise", "where", "worley"
]);
export const FUNCTION_META = {
@@ -54,6 +54,8 @@ export const FUNCTION_META = {
concatenate: { minArgs: 2, maxArgs: null, snippet: "concatenate()", description: "concatenate(x1, x2, ..., dim) - concatenates tensors along the specified dimension" },
conv: { minArgs: 2, maxArgs: null, snippet: "conv()", description: "convolution(tensor, [kernel_sizes...], kernel) - applies convolution with kernel. Expects [batch,channel, ...]" },
convolution: { minArgs: 2, maxArgs: null, snippet: "convolution()", description: "convolution(tensor, [kernel_sizes...], kernel) - applies convolution with kernel. Expects [batch,channel, ...]" },
coordinates: { minArgs: 2, maxArgs: 3, snippet: "coordinates()", description: "coords(shape, dim, [dtype]) - generates a tensor with the specified shape whose values are the coordinates of each element along the specified dimension. The dtype can be specified to control the data type of the output tensor." },
coords: { minArgs: 2, maxArgs: 3, snippet: "coords()", description: "coords(shape, dim, [dtype]) - generates a tensor with the specified shape whose values are the coordinates of each element along the specified dimension. The dtype can be specified to control the data type of the output tensor." },
corr: { minArgs: 2, maxArgs: 2, snippet: "corr()", description: "correlation(x, y) - computes the correlation between x and y" },
correlation: { minArgs: 2, maxArgs: 2, snippet: "correlation()", description: "correlation(x, y) - computes the correlation between x and y" },
cos: { minArgs: 1, maxArgs: 1, snippet: "cos()", description: "cos(x) - applies cosinus function to value or each element of value" },