sort - add desc+dim

This commit is contained in:
mcDandy
2026-08-28 14:55:09 +02:00
parent b7db90fc84
commit 9f2c766050
8 changed files with 4349 additions and 4278 deletions
+4 -4
View File
@@ -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
/**
+8 -2
View File
@@ -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
File diff suppressed because it is too large Load Diff
+3 -3
View File
@@ -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
File diff suppressed because it is too large Load Diff
+3 -3
View File
@@ -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" },