snorm allows dimension/s
This commit is contained in:
@@ -171,10 +171,6 @@ func1:
|
||||
tnorm(x) - normalises tensor or list by multiplication such that sum(x^2)==1.0 for each slice. The slice is last dimension of the tensor.
|
||||
*/
|
||||
| TNORM LPAREN expr RPAREN # TNormFunc
|
||||
/**
|
||||
snorm(x) - frobenius norm of tensor
|
||||
*/
|
||||
| SNORM LPAREN expr RPAREN # SNormFunc
|
||||
| FLOOR LPAREN expr RPAREN # FloorFunc
|
||||
| CEIL LPAREN expr RPAREN # CeilFunc
|
||||
| ROUND LPAREN expr RPAREN # RoundFunc
|
||||
@@ -241,6 +237,10 @@ func2:
|
||||
tmax(x,y) - elementwise maximum of tensor or list. max(x,y) when inputs are floats
|
||||
*/
|
||||
| TMAX LPAREN expr COMMA expr RPAREN # TMaxFunc
|
||||
/**
|
||||
snorm(x,[dim]) - frobenius norm of tensor. Along dimension if dimension specified. Dimension can be a list
|
||||
*/
|
||||
| SNORM LPAREN expr (COMMA expr)? RPAREN # SNormFunc
|
||||
| STEP LPAREN expr COMMA expr RPAREN # StepFunc
|
||||
| TOPK LPAREN expr COMMA expr RPAREN # TopkFunc
|
||||
| BOTK LPAREN expr COMMA expr RPAREN # BotkFunc
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from re import S
|
||||
import time
|
||||
import numpy as np
|
||||
import os
|
||||
from numpy._core.multiarray import scalar
|
||||
import torch
|
||||
import math
|
||||
import inspect
|
||||
@@ -220,18 +222,22 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
return list_op(val)
|
||||
return val
|
||||
|
||||
def _to_int(self, x, ctx, context_name="operation"):
|
||||
def _to_int(self, x, ctx, context_name="operation", strict=False):
|
||||
"""Convert value to int, handling tensors and nested lists recursively"""
|
||||
if self._is_tensor(x):
|
||||
if x.numel() == 1:
|
||||
return int(x.item())
|
||||
else:
|
||||
elif strict:
|
||||
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: {context_name} expects scalar dimensions, got tensor with shape {x.shape}")
|
||||
else:
|
||||
return x.int()
|
||||
elif self._is_list(x):
|
||||
if len(x) == 1:
|
||||
return self._to_int(x[0], ctx, context_name)
|
||||
else:
|
||||
elif strict:
|
||||
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: {context_name} expects scalar dimensions, got list with {len(x)} elements")
|
||||
else:
|
||||
return [self._to_int(v, ctx, context_name) for v in x]
|
||||
else:
|
||||
return int(float(x))
|
||||
|
||||
@@ -615,9 +621,14 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
return 1.0 if val != 0 else 0.0
|
||||
|
||||
def visitSNormFunc(self, ctx):
|
||||
val = (yield ctx.expr())
|
||||
val = (yield ctx.expr(0))
|
||||
if len(ctx.expr()) > 1:
|
||||
dim = self._to_int((yield ctx.expr(1)), ctx, "s_norm dimension")
|
||||
if self._is_tensor(val):
|
||||
return F.normalize(val, p=2, dim=dim)
|
||||
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: s_norm with dimension argument only supports tensors")
|
||||
if self._is_tensor(val):
|
||||
res = torch.linalg.norm(val)
|
||||
res = torch.linalg.norm(val, dim=-1)
|
||||
if res.numel() == 1:
|
||||
return float(res.item())
|
||||
return res
|
||||
@@ -1074,7 +1085,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: reshape expects scalar dimensions, got tensor with shape {d.shape}. Did you mean to pass a shape list instead of data?")
|
||||
if self._is_list(d) and len(d) > 1:
|
||||
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: reshape expects scalar dimensions, got list with {len(d)} elements. Did you pass a data variable (like V) instead of a shape?")
|
||||
result.append(self._to_int(d, ctx, "reshape"))
|
||||
result.append(self._to_int(d, ctx, "reshape", strict=True))
|
||||
new_shape = result
|
||||
elif isinstance(new_shape, (int, float)):
|
||||
new_shape = [int(float(new_shape))]
|
||||
@@ -2062,11 +2073,6 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
return float(v.item())
|
||||
return float(v)
|
||||
|
||||
def _to_int(v):
|
||||
if self._is_tensor(v):
|
||||
return int(v.item())
|
||||
return int(v)
|
||||
|
||||
def _to_bool(v):
|
||||
if isinstance(v, bool):
|
||||
return v
|
||||
@@ -2078,9 +2084,10 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
|
||||
text = _to_str((yield ctx.expr(0)))
|
||||
font_name = _to_str((yield ctx.expr(1)))
|
||||
size = max(1, _to_int((yield ctx.expr(2))))
|
||||
max_width = _to_int((yield ctx.expr(3))) if len(ctx.expr()) > 3 else 0
|
||||
weight_val = _to_int((yield ctx.expr(4))) if len(ctx.expr()) > 4 else 400
|
||||
size = max(1, self._to_int((yield ctx.expr(2)), ctx,"text_image", strict=True))
|
||||
size = max(1, self._to_int((yield ctx.expr(2)), ctx,"text_image", strict=True))
|
||||
max_width = self._to_int((yield ctx.expr(3)), ctx,"text_image", strict=True) if len(ctx.expr()) > 3 else 0
|
||||
weight_val = self._to_int((yield ctx.expr(4)), ctx,"text_image", strict=True) if len(ctx.expr()) > 4 else 400
|
||||
angle = _to_float((yield ctx.expr(5))) if len(ctx.expr()) > 5 else 0.0
|
||||
spacing = max(0.1, _to_float((yield ctx.expr(6)))) if len(ctx.expr()) > 6 else 1.0
|
||||
is_italic = _to_bool((yield ctx.expr(7))) if len(ctx.expr()) > 7 else False
|
||||
@@ -2950,9 +2957,9 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
shape_val = yield ctx.indexExpr()
|
||||
|
||||
if self._is_list(shape_val):
|
||||
shape = [self._to_int(v, ctx, "tensor") for v in shape_val]
|
||||
shape = self._to_int(v, ctx, "tensor")
|
||||
elif self._is_tensor(shape_val):
|
||||
shape = [self._to_int(v, ctx, "tensor") for v in shape_val.flatten().tolist()]
|
||||
shape = self._to_int(v, ctx, "tensor").flatten().tolist()
|
||||
else:
|
||||
shape = [self._to_int(shape_val, ctx, "tensor")]
|
||||
|
||||
@@ -2963,7 +2970,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
dim = -1
|
||||
if len(ctx.expr()) > 1:
|
||||
dim_val = (yield ctx.expr(1))
|
||||
dim = self._to_int(dim_val, ctx, "softmax dim")
|
||||
dim = self._to_int(dim_val, ctx, "softmax dim", strict=True)
|
||||
return F.softmax(val, dim=dim)
|
||||
|
||||
def visitSoftminFunc(self, ctx):
|
||||
@@ -2971,7 +2978,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
dim = -1
|
||||
if len(ctx.expr()) > 1:
|
||||
dim_val = (yield ctx.expr(1))
|
||||
dim = self._to_int(dim_val, ctx, "softmin dim")
|
||||
dim = self._to_int(dim_val, ctx, "softmin dim", strict=True)
|
||||
return F.softmax(-val, dim=dim)
|
||||
|
||||
def visitArgminFunc(self, ctx):
|
||||
@@ -4170,7 +4177,7 @@ class UnifiedMathVisitor(MathExprVisitor):
|
||||
return torch.linspace(start, end, steps, device=self.device)
|
||||
|
||||
def visitLogspaceFunc(self, ctx):
|
||||
"""linspace(start, end, steps) - linearly spaced values"""
|
||||
"""logspace(start, end, steps, base) - logarithmically spaced values"""
|
||||
start_val = yield ctx.expr(0)
|
||||
end_val = yield ctx.expr(1)
|
||||
steps_val = yield ctx.expr(2)
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -782,15 +782,6 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#SNormFunc.
|
||||
def enterSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#SNormFunc.
|
||||
def exitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#FloorFunc.
|
||||
def enterFloorFunc(self, ctx:MathExprParser.FloorFuncContext):
|
||||
pass
|
||||
@@ -1340,6 +1331,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#SNormFunc.
|
||||
def enterSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#SNormFunc.
|
||||
def exitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#StepFunc.
|
||||
def enterStepFunc(self, ctx:MathExprParser.StepFuncContext):
|
||||
pass
|
||||
|
||||
+2018
-2001
File diff suppressed because it is too large
Load Diff
@@ -439,11 +439,6 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#SNormFunc.
|
||||
def visitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#FloorFunc.
|
||||
def visitFloorFunc(self, ctx:MathExprParser.FloorFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -749,6 +744,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#SNormFunc.
|
||||
def visitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#StepFunc.
|
||||
def visitStepFunc(self, ctx:MathExprParser.StepFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
@@ -206,7 +206,7 @@ INBUILT_FUNCTION_META = {
|
||||
'smin': {'min_args': 1, 'max_args': None, 'snippet': 'smin()', 'description': ''},
|
||||
'smootherstep': {'min_args': 3, 'max_args': 3, 'snippet': 'smootherstep()', 'description': ''},
|
||||
'smoothstep': {'min_args': 3, 'max_args': 3, 'snippet': 'smoothstep()', 'description': ''},
|
||||
'snorm': {'min_args': 1, 'max_args': 1, 'snippet': 'snorm()', 'description': 'snorm(x) - frobenius norm of tensor'},
|
||||
'snorm': {'min_args': 1, 'max_args': 2, 'snippet': 'snorm()', 'description': 'snorm(x,[dim]) - frobenius norm of tensor. Along dimension if dimension specified. Dimension can be a list'},
|
||||
'softmax': {'min_args': 1, 'max_args': 2, 'snippet': 'softmax()', 'description': ''},
|
||||
'softmin': {'min_args': 1, 'max_args': 2, 'snippet': 'softmin()', 'description': ''},
|
||||
'softplus': {'min_args': 1, 'max_args': 1, 'snippet': 'softplus()', 'description': ''},
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -782,15 +782,6 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#SNormFunc.
|
||||
def enterSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#SNormFunc.
|
||||
def exitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#FloorFunc.
|
||||
def enterFloorFunc(self, ctx:MathExprParser.FloorFuncContext):
|
||||
pass
|
||||
@@ -1340,6 +1331,15 @@ class MathExprListener(ParseTreeListener):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#SNormFunc.
|
||||
def enterSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
pass
|
||||
|
||||
# Exit a parse tree produced by MathExprParser#SNormFunc.
|
||||
def exitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
pass
|
||||
|
||||
|
||||
# Enter a parse tree produced by MathExprParser#StepFunc.
|
||||
def enterStepFunc(self, ctx:MathExprParser.StepFuncContext):
|
||||
pass
|
||||
|
||||
+2236
-2218
File diff suppressed because it is too large
Load Diff
@@ -439,11 +439,6 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#SNormFunc.
|
||||
def visitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#FloorFunc.
|
||||
def visitFloorFunc(self, ctx:MathExprParser.FloorFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
@@ -749,6 +744,11 @@ class MathExprVisitor(ParseTreeVisitor):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#SNormFunc.
|
||||
def visitSNormFunc(self, ctx:MathExprParser.SNormFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
|
||||
# Visit a parse tree produced by MathExprParser#StepFunc.
|
||||
def visitStepFunc(self, ctx:MathExprParser.StepFuncContext):
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
@@ -209,7 +209,7 @@ export const FUNCTION_META = {
|
||||
smin: { minArgs: 1, maxArgs: null, snippet: "smin()", description: "" },
|
||||
smootherstep: { minArgs: 3, maxArgs: 3, snippet: "smootherstep()", description: "" },
|
||||
smoothstep: { minArgs: 3, maxArgs: 3, snippet: "smoothstep()", description: "" },
|
||||
snorm: { minArgs: 1, maxArgs: 1, snippet: "snorm()", description: "snorm(x) - frobenius norm of tensor" },
|
||||
snorm: { minArgs: 1, maxArgs: 2, snippet: "snorm()", description: "snorm(x,[dim]) - frobenius norm of tensor. Along dimension if dimension specified. Dimension can be a list" },
|
||||
softmax: { minArgs: 1, maxArgs: 2, snippet: "softmax()", description: "" },
|
||||
softmin: { minArgs: 1, maxArgs: 2, snippet: "softmin()", description: "" },
|
||||
softplus: { minArgs: 1, maxArgs: 1, snippet: "softplus()", description: "" },
|
||||
|
||||
Reference in New Issue
Block a user