roll linspace, logspace - no grammer build
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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';
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user