diff --git a/README.md b/README.md index 35f4503..f070f2c 100644 --- a/README.md +++ b/README.md @@ -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) diff --git a/more_math/Parser/MathExpr.g4 b/more_math/Parser/MathExpr.g4 index ab0314c..f8ab693 100644 --- a/more_math/Parser/MathExpr.g4 +++ b/more_math/Parser/MathExpr.g4 @@ -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'; diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index 3b56e3f..99cb6e0 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -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) + + \ No newline at end of file diff --git a/web/script_text_input.js b/web/script_text_input.js index 4fcd0da..f9de1d8 100644 --- a/web/script_text_input.js +++ b/web/script_text_input.js @@ -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",