snorm allows dimension/s

This commit is contained in:
mcDandy
2026-06-15 21:49:26 +02:00
parent 5ac9fbb62c
commit 2ebe252045
12 changed files with 4316 additions and 4274 deletions
+4 -4
View File
@@ -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
+26 -19
View File
@@ -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
+9 -9
View File
@@ -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
File diff suppressed because it is too large Load Diff
+5 -5
View File
@@ -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)
+1 -1
View File
@@ -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
+9 -9
View File
@@ -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
File diff suppressed because it is too large Load Diff
+5 -5
View File
@@ -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)
+1 -1
View File
@@ -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: "" },