This commit is contained in:
mcDandy
2026-01-10 00:00:34 +01:00
parent 08cb017e57
commit 4e744a8e82
+34 -44
View File
@@ -139,51 +139,46 @@ class UnifiedMathVisitor(MathExprVisitor):
return self._bool_op(self.visit(ctx.compExpr()), self.visit(ctx.addExpr()), torch.le, lambda a, b: float(a <= b))
# Functions
def _func_dispatch(self, arg, torch_fn, scalar_fn):
if self._is_list(arg):
return [self._func_dispatch(x, torch_fn, scalar_fn) for x in arg]
if self._is_tensor(arg):
return torch_fn(arg).contiguous()
return scalar_fn(arg)
def visitSinFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.sin, math.sin)
return self._unary_op(self.visit(ctx.expr()), torch.sin, math.sin)
def visitCosFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.cos, math.cos)
return self._unary_op(self.visit(ctx.expr()), torch.cos, math.cos)
def visitTanFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.tan, math.tan)
return self._unary_op(self.visit(ctx.expr()), torch.tan, math.tan)
def visitAsinFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.asin, math.asin)
return self._unary_op(self.visit(ctx.expr()), torch.asin, math.asin)
def visitAcosFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.acos, math.acos)
return self._unary_op(self.visit(ctx.expr()), torch.acos, math.acos)
def visitAtanFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.atan, math.atan)
return self._unary_op(self.visit(ctx.expr()), torch.atan, math.atan)
def visitSinhFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.sinh, math.sinh)
return self._unary_op(self.visit(ctx.expr()), torch.sinh, math.sinh)
def visitCoshFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.cosh, math.cosh)
return self._unary_op(self.visit(ctx.expr()), torch.cosh, math.cosh)
def visitTanhFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.tanh, math.tanh)
return self._unary_op(self.visit(ctx.expr()), torch.tanh, math.tanh)
def visitAsinhFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.asinh, math.asinh)
return self._unary_op(self.visit(ctx.expr()), torch.asinh, math.asinh)
def visitAcoshFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.acosh, math.acosh)
return self._unary_op(self.visit(ctx.expr()), torch.acosh, math.acosh)
def visitAtanhFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.atanh, math.atanh)
return self._unary_op(self.visit(ctx.expr()), torch.atanh, math.atanh)
def visitAbsFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.abs, abs)
return self._unary_op(self.visit(ctx.expr()), torch.abs, abs)
def visitAbsExp(self, ctx):
val = self.visit(ctx.expr())
@@ -194,59 +189,54 @@ class UnifiedMathVisitor(MathExprVisitor):
return abs(val)
def visitSqrtFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.sqrt, math.sqrt)
return self._unary_op(self.visit(ctx.expr()), torch.sqrt, math.sqrt)
def visitLnFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.log, math.log)
return self._unary_op(self.visit(ctx.expr()), torch.log, math.log)
def visitLogFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.log10, math.log10)
return self._unary_op(self.visit(ctx.expr()), torch.log10, math.log10)
def visitExpFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.exp, math.exp)
return self._unary_op(self.visit(ctx.expr()), torch.exp, math.exp)
def visitFloorFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.floor, math.floor)
return self._unary_op(self.visit(ctx.expr()), torch.floor, math.floor)
def visitCeilFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.ceil, math.ceil)
return self._unary_op(self.visit(ctx.expr()), torch.ceil, math.ceil)
def visitRoundFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.round, round)
return self._unary_op(self.visit(ctx.expr()), torch.round, round)
def visitSignFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.sign, lambda x: (1.0 if x > 0 else (-1.0 if x < 0 else 0.0)))
return self._unary_op(self.visit(ctx.expr()), torch.sign, lambda x: (1.0 if x > 0 else (-1.0 if x < 0 else 0.0)))
def visitFractFunc(self, ctx):
val = self.visit(ctx.expr())
if self._is_list(val):
return [x - math.floor(x) for x in val]
if self._is_tensor(val):
return val - torch.floor(val)
return val - math.floor(val)
return self._unary_op(self.visit(ctx.expr()), lambda x: x - torch.floor(x), lambda x: x - math.floor(x))
def visitGammaFunc(self, ctx):
torch_gamma = getattr(torch.special, "gamma", None)
if torch_gamma is None:
torch_gamma = lambda x: torch.exp(torch.lgamma(x))
return self._func_dispatch(self.visit(ctx.expr()), torch_gamma, math.gamma)
return self._unary_op(self.visit(ctx.expr()), torch_gamma, math.gamma)
def visitSigmoidFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.sigmoid, lambda x: 1.0 / (1.0 + math.exp(-x)))
return self._unary_op(self.visit(ctx.expr()), torch.sigmoid, lambda x: 1.0 / (1.0 + math.exp(-x)))
def visitReluFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.relu, lambda x: max(0.0, x))
return self._unary_op(self.visit(ctx.expr()), torch.relu, lambda x: max(0.0, x))
def visitSoftplusFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), F.softplus, lambda x: math.log(1.0 + math.exp(x)))
return self._unary_op(self.visit(ctx.expr()), F.softplus, lambda x: math.log(1.0 + math.exp(x)))
def visitGeluFunc(self, ctx):
return self._func_dispatch(
return self._unary_op(
self.visit(ctx.expr()), F.gelu, lambda x: 0.5 * x * (1 + math.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * math.pow(x, 3))))
)
def visitAnglFunc(self, ctx):
return self._func_dispatch(self.visit(ctx.expr()), torch.angle, lambda x: math.atan2(0, x) if x < 0 else 0)
return self._unary_op(self.visit(ctx.expr()), torch.angle, lambda x: math.atan2(0, x) if x < 0 else 0)
def visitPrintFunc(self, ctx):
val = self.visit(ctx.expr())
@@ -799,7 +789,7 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitConvFunc(self, ctx):
tensor, kernel_val, kernel_sizes, spatial_dims_count = self._parse_conv_args(ctx)
# Expect (Batch..., Channel, Spatial...)
# spatial_dims_count = 1, 2, or 3
# Must have at least Channel + Spatial dims
@@ -810,13 +800,13 @@ class UnifiedMathVisitor(MathExprVisitor):
spatial_shape = tensor.shape[-spatial_dims_count:]
in_channels = tensor.shape[-(spatial_dims_count + 1)]
batch_shape = tensor.shape[:-(spatial_dims_count + 1)]
total_batch = 1
for s in batch_shape:
total_batch *= s
conv_input = tensor.reshape(total_batch, in_channels, *spatial_shape)
result = self._apply_conv_internal(conv_input, kernel_val, kernel_sizes, spatial_dims_count)
return result.reshape(*batch_shape, in_channels, *spatial_shape)