Files
mcDandy-more_math/more_math/Parser/TensorEvalVisitor.py
T
mcDandy e44fd922d4 added map,conv and permute
map - maps points based on equations
conv - allows running convolution filter
permute - switches dimension position
2025-12-26 18:42:47 +01:00

558 lines
22 KiB
Python

import torch
import torch.special
from .MathExprVisitor import MathExprVisitor
class TensorEvalVisitor(MathExprVisitor):
def __init__(self, variables, shape, device=None):
self.variables = variables
self.spatial_variables = variables.copy()
self.shape = shape
# Infer device from variables if not provided
if device is None:
self.device = next((v.device for v in variables.values() if isinstance(v, torch.Tensor)), torch.device("cpu"))
else:
self.device = device
def _fold_nd(self, tsr, spatial_dims):
# Ensure tensor has rank (spatial_dims + 2)
# 1. Unsqueeze if too low rank (add dummy channel at dim 1)
original_shape = tsr.shape
added_dims = 0
target_rank = spatial_dims + 2
while tsr.ndim < target_rank:
tsr = tsr.unsqueeze(1)
added_dims += 1
# 2. Fold leading dimensions into batch if too high rank
folded = False
if tsr.ndim > target_rank:
fold_count = tsr.ndim - target_rank
new_batch = 1
for i in range(fold_count + 1):
new_batch *= tsr.shape[i]
tsr = tsr.reshape(new_batch, *tsr.shape[fold_count+1:])
folded = True
else:
folded = (added_dims > 0)
return tsr, original_shape, added_dims, folded
def _unfold_nd(self, tsr, original_shape, added_dims, folded):
spatial_dims = tsr.ndim - 2
# Restore folded batch if any
if folded and added_dims == 0:
target_fold_rank = len(original_shape) - (spatial_dims + 1)
fold_dims = original_shape[:target_fold_rank]
tsr = tsr.reshape(*fold_dims, *tsr.shape[1:])
# Squeeze added dims (dummy channels)
for _ in range(added_dims):
if tsr.ndim > len(original_shape) and tsr.shape[1] == 1:
tsr = tsr.squeeze(1)
return tsr
def visitPermuteFunc(self, ctx):
tsr = self.visit(ctx.expr(0))
# permute(tsr, [d1, d2, ...]) or permute(tsr, d1, d2, ...)
# Actually FuncN grammar was not used for permute but it could be.
# Let's check the grammar. I added `PERM '(' expr ',' expr ')'` which is Func2 style.
# But for permute we need a list.
dims = self.visit(ctx.expr(1))
if isinstance(dims, torch.Tensor):
dims = dims.flatten().long().tolist()
return tsr.permute(*dims)
def visitReshapeFunc(self, ctx):
tsr = self.visit(ctx.expr(0))
new_shape = self.visit(ctx.expr(1))
if isinstance(new_shape, torch.Tensor):
new_shape = new_shape.flatten().long().tolist()
return tsr.reshape(*new_shape)
def visitNumberExp(self, ctx):
return torch.tensor(float(ctx.getText()), device=self.device)
def visitConstantExp(self, ctx):
name = ctx.getText().lower()
if name == "pi":
return torch.full(self.shape, 3.141592653589793, device=self.device)
if name == "e":
return torch.full(self.shape, 2.718281828459045, device=self.device)
raise ValueError(f"Unknown constant: {name}")
def visitVariableExp(self, ctx):
name = ctx.getText()
if name not in self.variables:
raise ValueError(f"Variable '{name}' not found")
val = self.variables[name]
if not isinstance(val, torch.Tensor):
return torch.tensor(float(val), device=self.device)
return val
def visitParenExp(self, ctx):
return self.visit(ctx.expr())
def visitUnaryPlus(self, ctx):
return +self.visit(ctx.unaryExpr())
def visitUnaryMinus(self, ctx):
return -self.visit(ctx.unaryExpr())
def visitAddExp(self, ctx):
return torch.add(self.visit(ctx.addExpr()), self.visit(ctx.mulExpr()))
def visitSubExp(self, ctx):
return torch.sub(self.visit(ctx.addExpr()), self.visit(ctx.mulExpr()))
def visitMulExp(self, ctx):
return torch.mul(self.visit(ctx.mulExpr()), self.visit(ctx.powExpr()))
def visitDivExp(self, ctx):
return torch.div(self.visit(ctx.mulExpr()), self.visit(ctx.powExpr()))
def visitModExp(self, ctx):
return torch.fmod(self.visit(ctx.mulExpr()), self.visit(ctx.powExpr()))
def visitPowExp(self, ctx):
return torch.pow(self.visit(ctx.unaryExpr()), self.visit(ctx.powExpr()))
def visitToUnary(self, ctx):
return self.visit(ctx.unaryExpr())
def visitToPow(self, ctx):
return self.visit(ctx.powExpr())
def visitToMul(self, ctx):
return self.visit(ctx.mulExpr())
def visitToAdd(self, ctx):
return self.visit(ctx.addExpr())
def visitToAnd(self, ctx):
return self.visit(ctx.andExpr())
def visitToXor(self, ctx):
return self.visit(ctx.xorExpr())
def visitOrExp(self, ctx):
return torch.logical_or(self.visit(ctx.orExpr()).bool(), self.visit(ctx.xorExpr()).bool())
def visitXorExp(self, ctx):
return torch.pow(self.visit(ctx.xorExpr()), self.visit(ctx.andExpr()))
def visitAndExp(self, ctx):
return torch.logical_and(self.visit(ctx.andExpr()).bool(), self.visit(ctx.addExpr()).bool())
def visitNeExp(self, ctx):
return torch.ne(self.visit(ctx.compExpr()), self.visit(ctx.addExpr())).float()
def visitEqExp(self, ctx):
return torch.eq(self.visit(ctx.compExpr()), self.visit(ctx.addExpr())).float()
def visitGtExp(self, ctx):
return torch.gt(self.visit(ctx.compExpr()), self.visit(ctx.addExpr())).float()
def visitLtExp(self, ctx):
return torch.lt(self.visit(ctx.compExpr()), self.visit(ctx.addExpr())).float()
def visitGeExp(self, ctx):
return torch.ge(self.visit(ctx.compExpr()), self.visit(ctx.addExpr())).float()
def visitLeExp(self, ctx):
return torch.le(self.visit(ctx.compExpr()), self.visit(ctx.addExpr())).float()
# Single-argument functions
def visitSinFunc(self, ctx): return torch.sin(self.visit(ctx.expr()))
def visitCosFunc(self, ctx): return torch.cos(self.visit(ctx.expr()))
def visitTanFunc(self, ctx): return torch.tan(self.visit(ctx.expr()))
def visitAsinFunc(self, ctx): return torch.asin(self.visit(ctx.expr()))
def visitAcosFunc(self, ctx): return torch.acos(self.visit(ctx.expr()))
def visitAtanFunc(self, ctx): return torch.atan(self.visit(ctx.expr()))
def visitSinhFunc(self, ctx): return torch.sinh(self.visit(ctx.expr()))
def visitCoshFunc(self, ctx): return torch.cosh(self.visit(ctx.expr()))
def visitTanhFunc(self, ctx): return torch.tanh(self.visit(ctx.expr()))
def visitAsinhFunc(self, ctx): return torch.asinh(self.visit(ctx.expr()))
def visitAcoshFunc(self, ctx): return torch.acosh(self.visit(ctx.expr()))
def visitAtanhFunc(self, ctx): return torch.atanh(self.visit(ctx.expr()))
def visitAbsFunc(self, ctx): return torch.abs(self.visit(ctx.expr()))
def visitAbsExp(self, ctx): return torch.abs(self.visit(ctx.expr()))
def visitListExp(self, ctx):
vals = [self.visit(e) for e in ctx.expr()]
# If any are tensors with shape, stack them or create a tensor
if all(v.ndim == 0 for v in vals):
return torch.tensor([v.item() for v in vals], device=self.device)
else:
# Broadcasting stack
return torch.stack(torch.broadcast_tensors(*vals), dim=0)
def visitNormExp(self, ctx): return torch.linalg.norm(self.visit(ctx.expr()))
def visitSqrtFunc(self, ctx): return torch.sqrt(self.visit(ctx.expr()))
def visitLnFunc(self, ctx): return torch.log(self.visit(ctx.expr()))
def visitLogFunc(self, ctx): return torch.log10(self.visit(ctx.expr()))
def visitExpFunc(self, ctx): return torch.exp(self.visit(ctx.expr()))
def visitTNormFunc(self, ctx): return torch.nn.functional.normalize(self.visit(ctx.expr()), p=2, dim=-1)
def visitSNormFunc(self, ctx): return torch.full(self.shape, torch.linalg.norm(self.visit(ctx.expr())).data[0], device=self.device)
def visitFloorFunc(self, ctx): return torch.floor(self.visit(ctx.expr()))
def visitCeilFunc(self, ctx): return torch.ceil(self.visit(ctx.expr()))
def visitRoundFunc(self, ctx): return torch.round(self.visit(ctx.expr()))
def visitGammaFunc(self, ctx): return torch.special.gamma(self.visit(ctx.expr())).exp()
def visitSigmoidFunc(self, ctx): return torch.sigmoid(self.visit(ctx.expr()))
def visitReluFunc(self, ctx): return torch.relu(self.visit(ctx.expr()))
def visitSoftplusFunc(self, ctx): return torch.nn.functional.softplus(self.visit(ctx.expr()))
def visitGeluFunc(self, ctx): return torch.nn.functional.gelu(self.visit(ctx.expr()))
def visitSignFunc(self, ctx): return torch.sign(self.visit(ctx.expr()))
def visitFractFunc(self, ctx):
val = self.visit(ctx.expr())
return val - torch.floor(val)
def visitAnglFunc(self, ctx): return torch.angle(self.visit(ctx.expr()))
def visitPrintFunc(self, ctx):
val = self.visit(ctx.expr())
print(val,end="\n")
return val
def visitSfftFunc(self, ctx):
"""
Spatial FFT - transforms the expression to frequency domain.
Applies FFT on all dimensions of the tensor.
"""
old_vars = self.variables
self.variables = self.spatial_variables.copy()
try:
val = self.visit(ctx.expr())
# Apply FFT on all dimensions
dims = tuple(range(val.ndim))
return torch.fft.fftn(val, dim=dims)
finally:
self.variables = old_vars
def visitSwapFunc(self, ctx):
tsr = self.visit(ctx.expr(0))
dim_t = self.visit(ctx.expr(1))
idx1_t = self.visit(ctx.expr(2))
idx2_t = self.visit(ctx.expr(3))
dim = int(dim_t.flatten()[0].item())
i = int(idx1_t.flatten()[0].item())
j = int(idx2_t.flatten()[0].item())
# Handle negative dim
while dim < 0: dim += tsr.ndim
while i < 0: i += tsr.shape[dim]
while j < 0: j += tsr.shape[dim]
indices = torch.arange(tsr.shape[dim], device=tsr.device)
val_i = indices[i].clone()
indices[i] = indices[j]
indices[j] = val_i
return torch.index_select(tsr, dim, indices)
def visitSifftFunc(self, ctx):
"""
Spatial IFFT - evaluates expression in frequency domain then transforms back.
Provides frequency coordinate variables for each dimension:
- Kx, Ky, Kz: frequency indices for last 3 dims (0-indexed)
- K: isotropic frequency magnitude (Euclidean distance from DC)
- Fx, Fy, Fz: size of each frequency dimension
"""
old_vars = self.variables
self.variables = self.variables.copy()
device = self.device
ndim = len(self.shape)
dim_names = ['x', 'y', 'z', 'w', 'v', 'u'] # Names for dims (from last to first)
# Generate frequency coordinates for each dimension
k_components = []
for i in range(ndim):
dim_idx = ndim - 1 - i # Start from last dim
size_d = self.shape[dim_idx]
# Create frequency indices (0-indexed, will be shifted for DC centering if needed)
values = torch.arange(size_d, dtype=torch.float32, device=device)
view_shape = [1] * ndim
view_shape[dim_idx] = size_d
values = values.view(*view_shape).expand(*self.shape)
# Bind named variables (Kx, Ky, Kz for last 3 dims)
if i < len(dim_names):
var_name = f'K{dim_names[i]}'
self.variables[var_name] = values
self.variables[f'F{dim_names[i]}'] = float(size_d)
# Generic fallback for all dims
self.variables[f'K_dim{dim_idx}'] = values
self.variables[f'F_dim{dim_idx}'] = float(size_d)
k_components.append(values)
# Calculate isotropic K (Euclidean distance from DC)
k_sq_sum = torch.zeros(self.shape, device=device)
for k_val in k_components:
k_sq_sum = k_sq_sum + k_val ** 2
self.variables['K'] = torch.sqrt(k_sq_sum)
self.variables['frequency'] = self.variables['K']
# Legacy aliases
if 'Kx' in self.variables:
self.variables['frequency_count'] = self.variables.get('Fx', 1.0)
try:
val = self.visit(ctx.expr())
# Apply IFFT on all dimensions
dims = tuple(range(val.ndim))
return torch.fft.ifftn(val, dim=dims).real
finally:
self.variables = old_vars
# Two-argument functions
def visitPowFunc(self, ctx):
return torch.pow(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)))
def visitAtan2Func(self, ctx):
return torch.atan2(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)))
def visitClampFunc(self, ctx): return torch.clamp(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)), self.visit(ctx.expr(2)))
def visitLerpFunc(self, ctx):
a = self.visit(ctx.expr(0))
b = self.visit(ctx.expr(1))
w = self.visit(ctx.expr(2))
return torch.lerp(a, b, w)
def visitSmoothstepFunc(self, ctx):
x = self.visit(ctx.expr(0))
edge0 = self.visit(ctx.expr(1))
edge1 = self.visit(ctx.expr(2))
# Scale, bias and saturate x to 0..1 range
t = torch.clamp((x - edge0) / (edge1 - edge0), 0.0, 1.0)
return t * t * (3.0 - 2.0 * t)
def visitStepFunc(self, ctx):
x = self.visit(ctx.expr(0))
edge = self.visit(ctx.expr(1))
# step(edge, x) = 1 if x >= edge else 0
return torch.where(x >= edge, 1.0, 0.0)
# N-argument functions
def visitSMinFunc(self, ctx):
args = [self.visit(e) for e in ctx.expr()]
if len(args) == 1:
return torch.min(args[0]) # Global min of single tensor
return torch.min(torch.stack(torch.broadcast_tensors(*args)))
def visitSMaxFunc(self, ctx):
args = [self.visit(e) for e in ctx.expr()]
if len(args) == 1:
return torch.max(args[0]) # Global max of single tensor
return torch.max(torch.stack(torch.broadcast_tensors(*args)))
def visitTMinFunc(self, ctx):
return torch.minimum(self.visit(ctx.expr(0)),self.visit(ctx.expr(1)))
def visitTMaxFunc(self, ctx):
return torch.maximum(self.visit(ctx.expr(0)),self.visit(ctx.expr(1)))
def visitFunc1Exp(self, ctx):
return self.visitChildren(ctx)
def visitFunc2Exp(self, ctx):
return self.visitChildren(ctx)
def visitFuncNExp(self, ctx):
return self.visitChildren(ctx)
def visitFunc3Exp(self, ctx):
return self.visitChildren(ctx)
def visitFunc4Exp(self, ctx):
return self.visitChildren(ctx)
def _normalize_coord(self, coord, size):
if size > 1:
return (coord / (size - 1)) * 2.0 - 1.0
return torch.zeros_like(coord)
def visitMapFunc(self, ctx):
tensor = self.visit(ctx.expr(0))
coords = [self.visit(ctx.expr(i)) for i in range(1, len(ctx.expr()))]
num_coords = len(coords)
if num_coords == 0:
return tensor
is_1d = (num_coords == 1)
if is_1d:
tensor = tensor.unsqueeze(-2)
zeros = torch.zeros_like(coords[0])
coords.append(zeros)
working_num_coords = len(coords)
spatial_in_shape = tensor.shape[-working_num_coords:]
leading_shape = tensor.shape[:-working_num_coords]
batch_size = 1
for s in leading_shape:
batch_size *= s
if len(leading_shape) == 0:
input_view = tensor.view(1, 1, *spatial_in_shape)
else:
input_view = tensor.view(batch_size, 1, *spatial_in_shape)
norm_coords_list = []
for i, coord in enumerate(coords):
dim_size = spatial_in_shape[-(i+1)]
norm = self._normalize_coord(coord, dim_size)
norm_coords_list.append(norm)
try:
broadcasted_coords = torch.broadcast_tensors(*norm_coords_list)
except RuntimeError:
raise ValueError(f"map(): Coordinate shapes {[c.shape for c in coords]} cannot be broadcast together.")
grid = torch.stack(broadcasted_coords, dim=-1)
grid_spatial_shape = grid.shape[:-1]
is_batched = False
if len(leading_shape) > 0 and len(grid_spatial_shape) >= len(leading_shape):
if grid_spatial_shape[:len(leading_shape)] == leading_shape:
is_batched = True
if batch_size > 1 and not is_batched:
grid = grid.expand(batch_size, *grid.shape)
grid_spatial_shape = grid.shape[1:-1]
elif is_batched:
flatten_shape = (batch_size,) + grid_spatial_shape[len(leading_shape):] + (grid.shape[-1],)
grid = grid.reshape(flatten_shape)
grid_spatial_shape = flatten_shape[1:-1]
total_output_elements = grid.numel() // working_num_coords // batch_size
if working_num_coords == 2:
grid_view = grid.reshape(batch_size, 1, total_output_elements, 2)
output = torch.nn.functional.grid_sample(
input_view, grid_view, mode='bilinear', padding_mode='zeros', align_corners=True
)
elif working_num_coords == 3:
grid_view = grid.reshape(batch_size, 1, 1, total_output_elements, 3)
output = torch.nn.functional.grid_sample(
input_view, grid_view, mode='bilinear', padding_mode='zeros', align_corners=True
)
else:
raise ValueError(f"map() supports up to 3 coordinate dimensions, got {num_coords} original coords.")
output = output.view(batch_size, total_output_elements)
final_shape = list(leading_shape) + list(grid_spatial_shape)
output = output.view(final_shape)
return output
def _kernel_coords(self, size, device):
half = size // 2
return torch.arange(size, device=device).float() - half
def visitConvFunc(self, ctx):
"""
N-dimensional convolution.
conv(tensor, kw, expr) - 1D conv on last dim
conv(tensor, kw, kh, expr) - 2D conv on last 2 dims
conv(tensor, kw, kh, kd, expr) - 3D conv on last 3 dims
Kernel expression uses kX, kY, kZ as centered coordinates.
"""
img = self.visit(ctx.expr(0))
num_args = len(ctx.expr())
# Parse kernel sizes based on argument count
if num_args == 3:
kernel_sizes = [int(self.visit(ctx.expr(1)).flatten()[0].item())]
k_expr_ctx = ctx.expr(2)
elif num_args == 4:
kernel_sizes = [
int(self.visit(ctx.expr(1)).flatten()[0].item()),
int(self.visit(ctx.expr(2)).flatten()[0].item())
]
k_expr_ctx = ctx.expr(3)
elif num_args == 5:
kernel_sizes = [
int(self.visit(ctx.expr(1)).flatten()[0].item()),
int(self.visit(ctx.expr(2)).flatten()[0].item()),
int(self.visit(ctx.expr(3)).flatten()[0].item())
]
k_expr_ctx = ctx.expr(4)
else:
raise ValueError("conv() expects 3-5 arguments: (tensor, kw, [kh], [kd], kernel_expr|list)")
num_spatial = len(kernel_sizes)
kw = kernel_sizes[0]
kh = kernel_sizes[1] if num_spatial >= 2 else 1
kd = kernel_sizes[2] if num_spatial >= 3 else 1
kx = self._kernel_coords(kw, self.device).view(1, 1, kw).expand(kd, kh, kw)
ky = self._kernel_coords(kh, self.device).view(1, kh, 1).expand(kd, kh, kw)
kz = self._kernel_coords(kd, self.device).view(kd, 1, 1).expand(kd, kh, kw)
k_variables = self.variables.copy()
k_variables.update({
'kX': kx, 'kY': ky, 'kZ': kz,
'kW': float(kw), 'kH': float(kh), 'kD': float(kd)
})
k_visitor = TensorEvalVisitor(k_variables, (kd, kh, kw), device=self.device)
kernel = k_visitor.visit(k_expr_ctx)
if kernel.numel() == 1:
kernel = kernel.expand(kd, kh, kw).clone()
elif kernel.numel() == kw * kh * kd:
kernel = kernel.reshape(kd, kh, kw)
leading_shape = img.shape[:-num_spatial] if num_spatial < img.ndim else ()
spatial_shape = img.shape[-num_spatial:]
combined_leading = 1
for s in leading_shape:
combined_leading *= s
if len(leading_shape) == 0:
reshaped = img.unsqueeze(0).unsqueeze(0)
else:
reshaped = img.reshape(combined_leading, 1, *spatial_shape)
channels = reshaped.shape[1]
if num_spatial == 1:
kernel_1d = kernel.flatten()[:kw]
conv_kernel = kernel_1d.view(1, 1, kw).expand(channels, 1, kw)
output = torch.nn.functional.conv1d(reshaped, conv_kernel, padding=kw//2, groups=channels)
elif num_spatial == 2:
kernel_2d = kernel[0] if kd > 1 else kernel.reshape(kh, kw)
conv_kernel = kernel_2d.view(1, 1, kh, kw).expand(channels, 1, kh, kw)
output = torch.nn.functional.conv2d(reshaped, conv_kernel, padding=(kh//2, kw//2), groups=channels)
elif num_spatial == 3:
conv_kernel = kernel.view(1, 1, kd, kh, kw).expand(channels, 1, kd, kh, kw)
output = torch.nn.functional.conv3d(reshaped, conv_kernel, padding=(kd//2, kh//2, kw//2), groups=channels)
else:
raise ValueError(f"conv() supports up to 3 spatial dimensions, got {num_spatial}")
output = output.squeeze(1)
if len(leading_shape) == 0:
output = output.squeeze(0)
else:
final_shape = list(leading_shape) + list(output.shape[1:])
output = output.reshape(final_shape)
return output
def visitAtomExp(self, ctx):
return self.visitChildren(ctx)
def visitFunc2Expr(self, ctx):
return self.visit(ctx.getChild(0)) # forward to Atan2Func, PowFunc, etc.
def visitExpr(self, ctx):
return self.visitChildren(ctx)