sort - add desc+dim
This commit is contained in:
@@ -261,15 +261,15 @@ func1:
|
||||
*/
|
||||
| VAR LPAREN expr RPAREN # VarFunc
|
||||
/**
|
||||
sort(x) - returns a sorted version of x (If input is tensor, it sorts the last dimension)
|
||||
sort(x, [desc], [dim]) - returns a sorted version of x (If input is tensor, it sorts the last dimension)
|
||||
*/
|
||||
| SORT LPAREN expr RPAREN # SortFunc
|
||||
| SORT LPAREN expr (COMMA expr (COMMA expr)?)? RPAREN # SortFunc
|
||||
/**
|
||||
any(x) - returns true if any element of x is non-zero
|
||||
any(x) - returns 1 if any element of x is non-zero otherwise 0
|
||||
*/
|
||||
| ANY LPAREN expr RPAREN # AnyFunc
|
||||
/**
|
||||
all(x) - returns true if all elements of x are non-zero
|
||||
all(x) - returns 1 if all elements of x are non-zero otherwise 0
|
||||
*/
|
||||
| ALL LPAREN expr RPAREN # AllFunc
|
||||
/**
|
||||
|
||||
@@ -1632,8 +1632,14 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
return float(torch.sum(self._bin_op(self._bin_op(x,a,torch.sub,lambda x, a: x - a,ctx),k,torch.pow,pow,ctx)).item())/x.numel()
|
||||
|
||||
def visitSortFunc(self, ctx):
|
||||
val = self._promote_to_tensor((yield ctx.expr()))
|
||||
sorted_val, _ = torch.sort(val)
|
||||
val = self._promote_to_tensor((yield ctx.expr(0)))
|
||||
desc = False
|
||||
dim = -1
|
||||
if len(ctx.expr()) > 1:
|
||||
desc = bool((yield ctx.expr(1)))
|
||||
if len(ctx.expr()) > 2:
|
||||
dim = int((yield ctx.expr(2)))
|
||||
sorted_val, _ = torch.sort(val, descending=desc, dim=dim)
|
||||
return sorted_val
|
||||
|
||||
def visitCossimFunc(self, ctx):
|
||||
|
||||
File diff suppressed because one or more lines are too long
+2126
-2094
File diff suppressed because it is too large
Load Diff
@@ -13,9 +13,9 @@ INBUILT_FUNCTION_META = {
|
||||
'abs': {'min_args': 1, 'max_args': 1, 'snippet': 'abs()', 'description': 'abs(x) - applies per element absolute value function. Same as |x| for numbers.'},
|
||||
'acos': {'min_args': 1, 'max_args': 1, 'snippet': 'acos()', 'description': 'acos(x) - applies arcus cosinus function to value or each element of value'},
|
||||
'acosh': {'min_args': 1, 'max_args': 1, 'snippet': 'acosh()', 'description': 'acosh(x) - applies hyperbolic arcus cosinus function to value or each element of value'},
|
||||
'all': {'min_args': 1, 'max_args': 1, 'snippet': 'all()', 'description': 'all(x) - returns true if all elements of x are non-zero'},
|
||||
'all': {'min_args': 1, 'max_args': 1, 'snippet': 'all()', 'description': 'all(x) - returns 1 if all elements of x are non-zero otherwise 0'},
|
||||
'angle': {'min_args': 1, 'max_args': 1, 'snippet': 'angle()', 'description': 'angle(x) - returns the angle of a complex number or vector'},
|
||||
'any': {'min_args': 1, 'max_args': 1, 'snippet': 'any()', 'description': 'any(x) - returns true if any element of x is non-zero'},
|
||||
'any': {'min_args': 1, 'max_args': 1, 'snippet': 'any()', 'description': 'any(x) - returns 1 if any element of x is non-zero otherwise 0'},
|
||||
'append': {'min_args': 2, 'max_args': 2, 'snippet': 'append()', 'description': 'append(x, y) - appends y to the end of x. If inputs are tensors use concatenate(x,...,dim)'},
|
||||
'argmax': {'min_args': 1, 'max_args': 2, 'snippet': 'argmax()', 'description': 'argmax(x, [as_position]) - returns the maximum position in flattened x or coordinates as a list when requested'},
|
||||
'argmin': {'min_args': 1, 'max_args': 2, 'snippet': 'argmin()', 'description': 'argmin(x, [as_position]) - returns the minimum position in flattened x or coordinates as a list when requested'},
|
||||
@@ -218,7 +218,7 @@ INBUILT_FUNCTION_META = {
|
||||
'softmax': {'min_args': 1, 'max_args': 2, 'snippet': 'softmax()', 'description': 'softmax(x, [axis]) - applies the softmax function to x (sum of slice == 1.0 and max value <= 1.0). axis defaults to last.'},
|
||||
'softmin': {'min_args': 1, 'max_args': 2, 'snippet': 'softmin()', 'description': 'softmin(x, [axis]) - applies the softmin function to x (same as softmax(-x,[axis])). axis defaults to last.'},
|
||||
'softplus': {'min_args': 1, 'max_args': 1, 'snippet': 'softplus()', 'description': 'softplus(x) - applies the softplus function: ln(1 + exp(x))'},
|
||||
'sort': {'min_args': 1, 'max_args': 1, 'snippet': 'sort()', 'description': 'sort(x) - returns a sorted version of x (If input is tensor, it sorts the last dimension)'},
|
||||
'sort': {'min_args': 2, 'max_args': 3, 'snippet': 'sort()', 'description': 'sort(x, [desc], [dim]) - returns a sorted version of x (If input is tensor, it sorts the last dimension)'},
|
||||
'split': {'min_args': 1, 'max_args': 2, 'snippet': 'split()', 'description': 'split(x, [delimiter]) - splits string to list of strings based on delimiter. Default is space.'},
|
||||
'sqrt': {'min_args': 1, 'max_args': 1, 'snippet': 'sqrt()', 'description': 'sqrt(x) - applies sqere root per element. Negative numbers return NaN (not a number)'},
|
||||
'squeeze': {'min_args': 1, 'max_args': 2, 'snippet': 'squeeze()', 'description': 'squeeze(x, [dim]) - removes size-1 dimensions from tensor x, optionally at dim'},
|
||||
|
||||
File diff suppressed because one or more lines are too long
+2203
-2170
File diff suppressed because it is too large
Load Diff
@@ -16,9 +16,9 @@ export const FUNCTION_META = {
|
||||
abs: { minArgs: 1, maxArgs: 1, snippet: "abs()", description: "abs(x) - applies per element absolute value function. Same as |x| for numbers." },
|
||||
acos: { minArgs: 1, maxArgs: 1, snippet: "acos()", description: "acos(x) - applies arcus cosinus function to value or each element of value" },
|
||||
acosh: { minArgs: 1, maxArgs: 1, snippet: "acosh()", description: "acosh(x) - applies hyperbolic arcus cosinus function to value or each element of value" },
|
||||
all: { minArgs: 1, maxArgs: 1, snippet: "all()", description: "all(x) - returns true if all elements of x are non-zero" },
|
||||
all: { minArgs: 1, maxArgs: 1, snippet: "all()", description: "all(x) - returns 1 if all elements of x are non-zero otherwise 0" },
|
||||
angle: { minArgs: 1, maxArgs: 1, snippet: "angle()", description: "angle(x) - returns the angle of a complex number or vector" },
|
||||
any: { minArgs: 1, maxArgs: 1, snippet: "any()", description: "any(x) - returns true if any element of x is non-zero" },
|
||||
any: { minArgs: 1, maxArgs: 1, snippet: "any()", description: "any(x) - returns 1 if any element of x is non-zero otherwise 0" },
|
||||
append: { minArgs: 2, maxArgs: 2, snippet: "append()", description: "append(x, y) - appends y to the end of x. If inputs are tensors use concatenate(x,...,dim)" },
|
||||
argmax: { minArgs: 1, maxArgs: 2, snippet: "argmax()", description: "argmax(x, [as_position]) - returns the maximum position in flattened x or coordinates as a list when requested" },
|
||||
argmin: { minArgs: 1, maxArgs: 2, snippet: "argmin()", description: "argmin(x, [as_position]) - returns the minimum position in flattened x or coordinates as a list when requested" },
|
||||
@@ -221,7 +221,7 @@ export const FUNCTION_META = {
|
||||
softmax: { minArgs: 1, maxArgs: 2, snippet: "softmax()", description: "softmax(x, [axis]) - applies the softmax function to x (sum of slice == 1.0 and max value <= 1.0). axis defaults to last." },
|
||||
softmin: { minArgs: 1, maxArgs: 2, snippet: "softmin()", description: "softmin(x, [axis]) - applies the softmin function to x (same as softmax(-x,[axis])). axis defaults to last." },
|
||||
softplus: { minArgs: 1, maxArgs: 1, snippet: "softplus()", description: "softplus(x) - applies the softplus function: ln(1 + exp(x))" },
|
||||
sort: { minArgs: 1, maxArgs: 1, snippet: "sort()", description: "sort(x) - returns a sorted version of x (If input is tensor, it sorts the last dimension)" },
|
||||
sort: { minArgs: 2, maxArgs: 3, snippet: "sort()", description: "sort(x, [desc], [dim]) - returns a sorted version of x (If input is tensor, it sorts the last dimension)" },
|
||||
split: { minArgs: 1, maxArgs: 2, snippet: "split()", description: "split(x, [delimiter]) - splits string to list of strings based on delimiter. Default is space." },
|
||||
sqrt: { minArgs: 1, maxArgs: 1, snippet: "sqrt()", description: "sqrt(x) - applies sqere root per element. Negative numbers return NaN (not a number)" },
|
||||
squeeze: { minArgs: 1, maxArgs: 2, snippet: "squeeze()", description: "squeeze(x, [dim]) - removes size-1 dimensions from tensor x, optionally at dim" },
|
||||
|
||||
Reference in New Issue
Block a user