roll linspace, logspace - no grammer build

This commit is contained in:
Daniel Martinek
2026-04-01 10:59:00 +02:00
parent 49ce366a98
commit 10be52d481
4 changed files with 55 additions and 4 deletions
+3
View File
@@ -159,6 +159,7 @@ You can also get the node from comfy manager under the name of More math.
- `overlay(base, overlay, offset)`: Replaces a region of `base` with `overlay` starting at `offset`. Works with strings (substring replacement), lists (element replacement), and tensors (region replacement). Areas outside the base are ignored.
- `pad(tensor, padding)`: Pads a tensor with specified padding (pair for each dimension). For example, `[1,2,0,0]` adds 1 element before and 2 elements after in the first dimension, and no padding in the second dimension.
- `concatenate(tensor1, tensor2, dim)` or `concat` or `cat`: Concatenates two tensors along specified dimension. It things of everything as tensor.
- `roll(tensor, shifts, dims)`: Rolls tensor along specified dimensions by given shifts. Elements that roll beyond the last position are re-introduced at the first position.
### Advanced Tensor Operations
@@ -220,6 +221,8 @@ You can also get the node from comfy manager under the name of More math.
- `print_shape(x)` or `pshp`: Prints the shape of x to the console and returns x.
- `pinv(x)`: Computes the permutation inverse of list. If `permute(i,x) = j`, then `permute(j,pinv(x)) = i`.
- `range(start, end, step)`: Generates a list of values from start (inclusive) to end (exclusive) with given step.
- `linspace(start, end, count)`: Generates a list of `count` evenly spaced values from start to end (inclusive).
- `logspace(start, end, count,base)`: Generates a list of `count` values logarithmically spaced between base^start and base^end.
- `nan_to_num(x, nan_value, posinf_value, neginf_value)` or `nvl`: Replaces NaN and infinite values in tensor with specified values.
- `remap(v, i_min, i_max, o_min, o_max)`: Remaps value `v` from input range `[i_min, i_max]` to output range `[o_min, o_max]`.
- `timestamp()` or `now`: Returns current UNIX timestamp (precision to microseconds, can be different on other systems)
+9 -2
View File
@@ -188,6 +188,7 @@ func2:
| COV LPAREN expr COMMA expr RPAREN # CovFunc
| CORR LPAREN expr COMMA expr RPAREN # CorrFunc
| APPEND LPAREN expr COMMA expr RPAREN # AppendFunc
| PERM LPAREN expr COMMA expr RPAREN # PermuteFunc
| GAUSSIAN LPAREN expr COMMA expr (COMMA expr)? RPAREN # GaussianFunc
| TOPK_IND LPAREN expr COMMA expr RPAREN # TopkIndFunc
| BOTK_IND LPAREN expr COMMA expr RPAREN # BotkIndFunc
@@ -226,12 +227,15 @@ func3:
| CROP LPAREN expr COMMA expr COMMA expr RPAREN # CropFunc
| SIFFT LPAREN expr (COMMA expr)? RPAREN # SifftFunc
| OVERLAY LPAREN expr COMMA expr COMMA expr RPAREN # OverlayFunc
| LINSPACE LPAREN expr COMMA expr COMMA expr RPAREN # LinspaceFunc
| ROLL LPAREN expr COMMA expr (COMMA expr)? RPAREN # RollFunc
| RGB_TO_HSV LPAREN expr (COMMA expr COMMA expr)? (COMMA expr)? RPAREN # RgbToHsvFunc
| HSV_TO_RGB LPAREN expr (COMMA expr COMMA expr)? (COMMA expr)? RPAREN # HsvToRgbFunc;
func4:
SWAP LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # SwapFunc
| NVL LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # NvlFunc
| LOGSPACE LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # ĹogspaceFunc
| DIST LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # DistFunc;
func5:
@@ -244,7 +248,6 @@ funcN:
| MAP LPAREN expr (COMMA expr)+ RPAREN # MapFunc
| EZCONV LPAREN expr (COMMA expr)+ RPAREN # EzConvFunc
| CONV LPAREN expr (COMMA expr)+ RPAREN # ConvFunc
| PERM LPAREN expr COMMA expr RPAREN # PermuteFunc
| RESHAPE LPAREN expr COMMA expr RPAREN # ReshapeFunc
| CONCAT LPAREN expr (COMMA expr)+ RPAREN # ConcatFunc;
@@ -315,12 +318,15 @@ SOFTPLUS: 'softplus';
GELU: 'gelu';
SIGN: 'sign';
MAP: 'map';
ROLL: 'roll';
EZCONV: 'ezconvolution' | 'ezconv';
CONV: 'convolution' | 'conv';
SWAP: 'swap';
PERM: 'permute' | 'perm';
RESHAPE: 'reshape' | 'rshp';
RANGE: 'range';
LINSPACE: 'linspace';
LOGSPACE: 'logspace';
TOPK: 'topk';
BOTK: 'botk';
PINV: 'pinv';
@@ -405,7 +411,7 @@ IN: 'in';
BREAK: 'break';
CONTINUE: 'continue';
RETURN: 'return';
TIMESTAMP: 'timestamp';
TIMESTAMP: 'timestamp'|'now';
SORT: 'sort';
ARGSORT: 'argsort';
ARGMIN: 'argmin';
@@ -414,6 +420,7 @@ SOFTMAX: 'softmax';
SOFTMIN: 'softmin';
UNIQUE: 'unique';
FLIP: 'flip';
ROLL: 'roll';
COV: 'cov';
CORR: 'corr' | 'correlation';
ENTROPY: 'entropy';
+41
View File
@@ -3472,3 +3472,44 @@ class UnifiedMathVisitor(MathExprVisitor):
b = (self._promote_to_tensor(b) * 256).clamp(0, 255).to(torch.int32)
return ((r << 16) | (g << 8) | b).contiguous()
return math.clamp(int(r*256),0,255) << 16 | math.clamp(int(g*256),0,255) << 8 | math.clamp(int(b*256),0,255)
def visitLinspaceFunc(self, ctx):
"""linspace(start, end, steps) - linearly spaced values"""
start_val = yield ctx.expr(0)
end_val = yield ctx.expr(1)
steps_val = yield ctx.expr(2)
start = float(start_val.item()) if self._is_tensor(start_val) else float(start_val)
end = float(end_val.item()) if self._is_tensor(end_val) else float(end_val)
steps = int(steps_val.item()) if self._is_tensor(steps_val) else int(steps_val)
return torch.linspace(start, end, steps, device=self.device)
def visitLogspaceFunc(self, ctx):
"""linspace(start, end, steps) - linearly spaced values"""
start_val = yield ctx.expr(0)
end_val = yield ctx.expr(1)
steps_val = yield ctx.expr(2)
base_val = yield ctx.expr(2)
start = float(start_val.item()) if self._is_tensor(start_val) else float(start_val)
end = float(end_val.item()) if self._is_tensor(end_val) else float(end_val)
base = float(base_val.item()) if self._is_tensor(base_val) else float(base_val)
steps = int(steps_val.item()) if self._is_tensor(steps_val) else int(steps_val)
return torch.logspace(start, end, steps, device=self.device)
def visitRollFunc(self, ctx):
"""roll(x, shift, [dim]) - circular shift of elements"""
x = self._promote_to_tensor((yield ctx.expr(0)))
shift_val = yield ctx.expr(1)
shift = int(shift_val.item()) if self._is_tensor(shift_val) else int(shift_val)
dim = 0
if len(ctx.expr()) > 2:
dim_val = yield ctx.expr(2)
dim = int(dim_val.item()) if self._is_tensor(dim_val) else int(dim_val)
return torch.roll(x, shifts=shift, dims=dim)
+2 -2
View File
@@ -18,9 +18,9 @@ const FUNCTIONS = new Set([
"distance", "remap", "cossim", "cosine_similarity", "count", "cnt", "length", "flatten", "append", "get_value",
"flow_apply", "batch_shuffle", "shuffle", "motion_mask", "flow_to_image", "overlay", "pad", "cross", "matmul", "rife",
"bnot", "bitwise_not", "bitcount", "popcount", "popcnt", "shape", "band", "bitwise_and", "bxor", "bitwise_xor",
"bor", "bitwise_or", "tensor", "stack_push", "stack_pop", "stack_clear", "stack_has", "stack_get", "timestamp",
"bor", "bitwise_or", "tensor", "stack_push", "stack_pop", "stack_clear", "stack_has", "stack_get", "timestamp","now",
"sort", "argsort", "argmin", "argmax", "softmax", "softmin", "unique", "flip", "cov", "corr", "correlation", "entropy",
"crop", "cat", "concatenate", "concat", "float","int",
"crop", "cat", "concatenate", "concat", "float","int", "linspace", "logspace", "roll",
"noise", "randn", "random_normal", "rand", "randu", "random_uniform", "randc", "random_cauchy", "rande",
"random_exponential", "randln", "random_log_normal", "randb", "random_bernoulli", "randp", "random_poisson", "randg",
"random_gamma", "randbeta", "random_beta", "randl", "random_laplace", "randgumbel", "random_gumbel", "randw",