diff --git a/more_math/Parser/FloatEvalVisitor.py b/more_math/Parser/FloatEvalVisitor.py deleted file mode 100644 index fbf6394..0000000 --- a/more_math/Parser/FloatEvalVisitor.py +++ /dev/null @@ -1,212 +0,0 @@ -import math -from .MathExprVisitor import MathExprVisitor - -class FloatEvalVisitor(MathExprVisitor): - def __init__(self, variables): - self.variables = variables - - def visitNumberExp(self, ctx): - print("Visiting number expression:", ctx.getText()) - return float(ctx.getText()) - - def visitConstantExp(self, ctx): - name = ctx.getText().lower() - if name == "pi": - return 3.141592653589793 - if name == "e": - return 2.718281828459045 - 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") - return self.variables[name] - - 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 self.visit(ctx.addExpr()) + self.visit(ctx.mulExpr()) - - def visitSubExp(self, ctx): - return self.visit(ctx.addExpr()) - self.visit(ctx.mulExpr()) - - def visitMulExp(self, ctx): - print("Visiting multiplication expression:", ctx.getText()) - if ctx.mulExpr() is None or ctx.powExpr() is None: - raise ValueError("Invalid multiplication expression") - return self.visit(ctx.mulExpr()) * self.visit(ctx.powExpr()) - - def visitDivExp(self, ctx): - return self.visit(ctx.mulExpr()) / self.visit(ctx.powExpr()) - - def visitModExp(self, ctx): - return self.visit(ctx.mulExpr()) % self.visit(ctx.powExpr()) - - def visitPowExp(self, ctx): - return math.pow(self.visit(ctx.unaryExpr()), self.visit(ctx.powExpr())) - - def visitNeExp(self, ctx): - return float(self.visit(ctx.compExpr()) != self.visit(ctx.addExpr())) - def visitEqExp(self, ctx): - return float(self.visit(ctx.compExpr()) == self.visit(ctx.addExpr())) - def visitGtExp(self, ctx): - return float(self.visit(ctx.compExpr()) > self.visit(ctx.addExpr())) - def visitLtExp(self, ctx): - return float(self.visit(ctx.compExpr()) < self.visit(ctx.addExpr())) - def visitGeExp(self, ctx): - return float(self.visit(ctx.compExpr()) >= self.visit(ctx.addExpr())) - def visitLeExp(self, ctx): - return float(self.visit(ctx.compExpr()) <= self.visit(ctx.addExpr())) - - 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 visitToGt(self, ctx): - return self.visit(ctx.gtExpr()) - def visitToLt(self, ctx): - return self.visit(ctx.ltExpr()) - def visitToEq(self, ctx): - return self.visit(ctx.eqExpr()) - def visitToNeq(self, ctx): - return self.visit(ctx.neqExpr()) - def visitToGe(self, ctx): - return self.visit(ctx.gteExpr()) - def visitToLe(self, ctx): - return self.visit(ctx.lteExpr()) - - # Single-argument functions - def visitSinFunc(self, ctx): return math.sin(self.visit(ctx.expr())) - def visitCosFunc(self, ctx): return math.cos(self.visit(ctx.expr())) - def visitTanFunc(self, ctx): return math.tan(self.visit(ctx.expr())) - def visitAsinFunc(self, ctx): return math.asin(self.visit(ctx.expr())) - def visitAcosFunc(self, ctx): return math.acos(self.visit(ctx.expr())) - def visitAtanFunc(self, ctx): return math.atan(self.visit(ctx.expr())) - def visitSinhFunc(self, ctx): return math.sinh(self.visit(ctx.expr())) - def visitCoshFunc(self, ctx): return math.cosh(self.visit(ctx.expr())) - def visitTanhFunc(self, ctx): return math.tanh(self.visit(ctx.expr())) - def visitAsinhFunc(self, ctx): return math.asinh(self.visit(ctx.expr())) - def visitAcoshFunc(self, ctx): return math.acosh(self.visit(ctx.expr())) - def visitAtanhFunc(self, ctx): return math.atanh(self.visit(ctx.expr())) - def visitAbsFunc(self, ctx): return abs(self.visit(ctx.expr())) - def visitAbsExp(self, ctx): return abs(self.visit(ctx.expr())) - def visitListExp(self, ctx): return self.visit(ctx.expr(0)) - def visitSqrtFunc(self, ctx): return math.sqrt(self.visit(ctx.expr())) - def visitLnFunc(self, ctx): return math.log(self.visit(ctx.expr())) - def visitLogFunc(self, ctx): return math.log10(self.visit(ctx.expr())) - def visitExpFunc(self, ctx): return math.exp(self.visit(ctx.expr())) - def visitNormFunc(self, ctx): - vals = self.visit(ctx.expr()) - if isinstance(vals, (list, tuple)): - return math.sqrt(sum(x**2 for x in vals) / len(vals)) - return abs(vals) - def visitFloorFunc(self, ctx): return math.floor(self.visit(ctx.expr())) - def visitFractFunc(self, ctx): - val = self.visit(ctx.expr()) - return val - math.floor(val) - def visitSigmoidFunc(self, ctx): return 1/(1+math.exp(-self.visit(ctx.expr()))) - def visitReluFunc(self, ctx): return max(0.0, self.visit(ctx.expr())) - def visitSoftplusFunc(self, ctx): - # log(1 + exp(x)) - x = self.visit(ctx.expr()) - # stability check? math.log1p(math.exp(x)) is better but might overflow for large x - if x > 20: return x - return math.log(1 + math.exp(x)) - def visitGeluFunc(self, ctx): - # 0.5 * x * (1 + erf(x / sqrt(2))) - x = self.visit(ctx.expr()) - return 0.5 * x * (1 + math.erf(x / 1.4142135623730951)) - def visitSignFunc(self, ctx): - x = self.visit(ctx.expr()) - return math.copysign(1.0, x) if x != 0 else 0.0 - def visitCeilFunc(self, ctx): return math.ceil(self.visit(ctx.expr())) - def visitRoundFunc(self, ctx): return round(self.visit(ctx.expr())) - def visitGammaFunc(self, ctx): return math.gamma(self.visit(ctx.expr())).exp() - def visitPrintFunc(self, ctx): - val = self.visit(ctx.expr()) - print(val,end="\n") - return val - - # Two-argument functions - def visitPowFunc(self, ctx): - return math.pow(self.visit(ctx.expr(0)), self.visit(ctx.expr(1))) - def visitAtan2Func(self, ctx): - return math.atan2(self.visit(ctx.expr(0)), self.visit(ctx.expr(1))) - def visitStepFunc(self, ctx): - x = self.visit(ctx.expr(0)) - edge = self.visit(ctx.expr(1)) - return 1.0 if x >= edge else 0.0 - - # N-argument functions - def visitSMinFunc(self, ctx): - args = [self.visit(e) for e in ctx.expr()] - return min(args) - def visitSMaxFunc(self, ctx): - args = [self.visit(e) for e in ctx.expr()] - return max(args) - - def visitClampFunc(self, ctx): - x = self.visit(ctx.expr(0)) - min_val = self.visit(ctx.expr(1)) - max_val = self.visit(ctx.expr(2)) - return max(min(x, max_val), min_val) - - def visitLerpFunc(self, ctx): - # a + (b - a) * w - a = self.visit(ctx.expr(0)) - b = self.visit(ctx.expr(1)) - w = self.visit(ctx.expr(2)) - return a + (b - a) * 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 = (x - edge0) / (edge1 - edge0) - t = max(0.0, min(1.0, t)) - # Evaluate polynomial - return t * t * (3.0 - 2.0 * t) - - def visitFunc1Exp(self, ctx): - return self.visitChildren(ctx) - def visitFunc2Exp(self, ctx): - return self.visitChildren(ctx) - def visitFunc3Exp(self, ctx): - return self.visitChildren(ctx) - def visitFunc4Exp(self, ctx): - return self.visitChildren(ctx) - - def visitMapFunc(self, ctx): - return self.visit(ctx.expr(0)) - - def visitConvFunc(self, ctx): - return self.visit(ctx.expr(0)) - - def visitFuncNExp(self, ctx): - return self.visitChildren(ctx) - 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) diff --git a/more_math/Parser/TensorEvalVisitor.py b/more_math/Parser/TensorEvalVisitor.py deleted file mode 100644 index 8ad8f43..0000000 --- a/more_math/Parser/TensorEvalVisitor.py +++ /dev/null @@ -1,538 +0,0 @@ -import torch -import torch.special - -from .MathExprVisitor import MathExprVisitor -from ..helper_functions import generate_dim_variables - -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.full(self.shape,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) - self.variables = self.variables | generate_dim_variables(values) - - 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 - if num_coords > 3: - raise ValueError("map() supports max 3 mapping functions.") - - spatial_in_shape = tensor.shape[-num_coords:] - leading_shape = tensor.shape[:-num_coords] - - batch_size = 1 - for s in leading_shape: batch_size *= s - - - input_view = tensor.reshape(batch_size, 1, *spatial_in_shape) - - norm_coords_list = [] - for i in range(num_coords): - dim_size = spatial_in_shape[i] - norm = self._normalize_coord(coords[i], dim_size) - norm_coords_list.append(norm) - - grid = torch.stack(norm_coords_list[::-1], dim=-1) - grid_spatial_shape = grid.shape[:-1] - - try: - grid_view = grid.reshape(batch_size, *grid_spatial_shape[-(num_coords):], num_coords) - except RuntimeError: - print("Reshape failed in map(); attempting expand workaround.") - grid_view = grid.expand(batch_size, *([-1] * len(grid_spatial_shape)), -1) - grid_view = grid_view.reshape(batch_size, *grid_view.shape[-(num_coords+1):-1], num_coords) - - if num_coords == 1: - input_final = input_view.reshape(batch_size, 1, 1, -1) - y_zeros = torch.zeros_like(grid_view[..., :1]) - grid_final = torch.cat([grid_view, y_zeros], dim=-1).unsqueeze(1) - output = torch.nn.functional.grid_sample(input_final, grid_final, align_corners=True) - - elif num_coords == 2: - grid_final = grid_view.reshape(batch_size, *grid_view.shape[-3:-1], 2) - output = torch.nn.functional.grid_sample(input_view, grid_final, align_corners=True) - else: - grid_final = grid_view.reshape(batch_size, *grid_view.shape[-4:-1], 3) - output = torch.nn.functional.grid_sample(input_view, grid_final, align_corners=True) - - actual_spatial = grid_view.shape[1:-1] - final_shape = list(leading_shape) + list(actual_spatial) - return output.reshape(final_shape) - - - 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), - 'kernel_width': float(kw), 'kernel_height': float(kh), 'kernel_depth': 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) - - def visitPrintShapeFunc(self, ctx): - tsr = self.visit(ctx.expr()) - print(tsr.shape) - return tsr \ No newline at end of file diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py new file mode 100644 index 0000000..ceaf87c --- /dev/null +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -0,0 +1,624 @@ +import torch +import math +import torch.nn.functional as F +from .MathExprVisitor import MathExprVisitor +from ..helper_functions import generate_dim_variables, getIndexTensorAlongDim + +class UnifiedMathVisitor(MathExprVisitor): + def __init__(self, variables, shape=None, device=None): + self.variables = variables + self.spatial_variables = variables.copy() + self.shape = shape if shape is not None else () + 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 _is_tensor(self, val): return isinstance(val, torch.Tensor) + def _is_list(self, val): return isinstance(val, (list, tuple)) + + def _promote_to_tensor(self, val): + if self._is_tensor(val): return val + if self._is_list(val): return torch.tensor(val, device=self.device) + return torch.tensor(val, device=self.device) + + def _bin_op(self, a, b, torch_op, scalar_op): + """ + Generic binary operation handler. + """ + # one of them is a list and one is tensor + if self._is_tensor(a) and self._is_list(b): return torch.stack([self._bin_op(a, x, torch_op, scalar_op) for x in b], dim=0) + if self._is_list(a) and self._is_tensor(b): return torch.stack([self._bin_op(x, b, torch_op, scalar_op) for x in a], dim=0) + + if self._is_list(a) and not self._is_tensor(b): + if self._is_list(b): + if len(a) != len(b): raise ValueError("List length mismatch") + return [self._bin_op(x, y, torch_op, scalar_op) for x, y in zip(a, b)] + return [self._bin_op(x, b, torch_op, scalar_op) for x in a] + + if not self._is_tensor(a) and self._is_list(b): + return [self._bin_op(a, x, torch_op, scalar_op) for x in b] + + if self._is_tensor(a) or self._is_tensor(b): + if torch_op: + return torch_op(a, b) + return scalar_op(a, b) + + return scalar_op(a, b) + + def _unary_op(self, a, torch_op, scalar_op): + if self._is_list(a): + return [self._unary_op(x, torch_op, scalar_op) for x in a] + if self._is_tensor(a): + return torch_op(a) if torch_op else scalar_op(a) + return scalar_op(a) + + # ======================== + # Visitors + # ======================== + + def visitNumberExp(self, ctx): + val_str = ctx.getText() + if '.' in val_str or 'e' in val_str: + return float(val_str) + return int(val_str) + + def visitConstantExp(self, ctx): + name = ctx.getText().lower() + if name == "pi": return math.pi + if name == "e": return math.e + 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") + return self.variables[name] + + def visitListExp(self, ctx): + return [self.visit(e) for e in ctx.expr()] + + def visitParenExp(self, ctx): + return self.visit(ctx.expr()) + + def visitUnaryPlus(self, ctx): + return self._unary_op(self.visit(ctx.unaryExpr()), lambda x: x, lambda x: +x) + + def visitUnaryMinus(self, ctx): + return self._unary_op(self.visit(ctx.unaryExpr()), torch.neg, lambda x: -x) + + # Binary Ops + def visitAddExp(self, ctx): + return self._bin_op(self.visit(ctx.addExpr()), self.visit(ctx.mulExpr()), + torch.add, lambda a,b: a+b) + + def visitSubExp(self, ctx): + return self._bin_op(self.visit(ctx.addExpr()), self.visit(ctx.mulExpr()), + torch.sub, lambda a,b: a-b) + + def visitMulExp(self, ctx): + return self._bin_op(self.visit(ctx.mulExpr()), self.visit(ctx.powExpr()), + torch.mul, lambda a,b: a*b) + + def visitDivExp(self, ctx): + return self._bin_op(self.visit(ctx.mulExpr()), self.visit(ctx.powExpr()), + torch.div, lambda a,b: a/b) + + def visitModExp(self, ctx): + return self._bin_op(self.visit(ctx.mulExpr()), self.visit(ctx.powExpr()), + torch.fmod, lambda a,b: a%b) + + def visitPowExp(self, ctx): + return self._bin_op(self.visit(ctx.unaryExpr()), self.visit(ctx.powExpr()), + torch.pow, math.pow) + + + def _bool_op(self, a, b, torch_op, scalar_op): + return self._bin_op(a, b, torch_op, scalar_op) + + def visitNeExp(self, ctx): return self._bool_op(self.visit(ctx.compExpr()), self.visit(ctx.addExpr()), torch.ne, lambda a,b: float(a!=b)) + def visitEqExp(self, ctx): return self._bool_op(self.visit(ctx.compExpr()), self.visit(ctx.addExpr()), torch.eq, lambda a,b: float(a==b)) + def visitGtExp(self, ctx): return self._bool_op(self.visit(ctx.compExpr()), self.visit(ctx.addExpr()), torch.gt, lambda a,b: float(a>b)) + def visitLtExp(self, ctx): return self._bool_op(self.visit(ctx.compExpr()), self.visit(ctx.addExpr()), torch.lt, lambda a,b: float(a=b)) + def visitLeExp(self, ctx): 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) + return scalar_fn(arg) + + def visitSinFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.sin, math.sin) + def visitCosFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.cos, math.cos) + def visitTanFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.tan, math.tan) + def visitAsinFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.asin, math.asin) + def visitAcosFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.acos, math.acos) + def visitAtanFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.atan, math.atan) + def visitSinhFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.sinh, math.sinh) + def visitCoshFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.cosh, math.cosh) + def visitTanhFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.tanh, math.tanh) + def visitAsinhFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.asinh, math.asinh) + def visitAcoshFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.acosh, math.acosh) + def visitAtanhFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.atanh, math.atanh) + + def visitAbsFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.abs, abs) + def visitAbsExp(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.abs, abs) + + def visitSqrtFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.sqrt, math.sqrt) + def visitLnFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.log, math.log) + def visitLogFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.log10, math.log10) + def visitExpFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.exp, math.exp) + + def visitFloorFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.floor, math.floor) + def visitCeilFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.ceil, math.ceil) + def visitRoundFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.round, round) + + def visitSignFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), torch.sign, lambda x: math.copysign(1.0, x)) + + 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) + + 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) + + def visitSigmoidFunc(self, ctx): return self._func_dispatch(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)) + def visitSoftplusFunc(self, ctx): return self._func_dispatch(self.visit(ctx.expr()), F.softplus, lambda x: math.log(1.0 + math.exp(x))) + def visitGeluFunc(self, ctx): return self._func_dispatch(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) + + def visitPrintFunc(self, ctx): + val = self.visit(ctx.expr()) + print(f"{val}") + return val + + def visitTNormFunc(self, ctx): + val = self.visit(ctx.expr()) + if self._is_tensor(val): return F.normalize(val, p=2, dim=-1) + return 1.0 if val != 0 else 0.0 + + def visitSNormFunc(self, ctx): + val = self.visit(ctx.expr()) + if self._is_tensor(val): return torch.linalg.norm(val) + return abs(val) + + # Two-argument functions + def visitPowFunc(self, ctx): return self._bin_op(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)), torch.pow, math.pow) + def visitAtan2Func(self, ctx): return self._bin_op(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)), torch.atan2, math.atan2) + def visitTMinFunc(self, ctx): return self._bin_op(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)), torch.minimum, min) + def visitTMaxFunc(self, ctx): return self._bin_op(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)), torch.maximum, max) + + def visitStepFunc(self, ctx): + # step(x, edge) = 1 if x >= edge else 0 + return self._bin_op(self.visit(ctx.expr(0)), self.visit(ctx.expr(1)), + lambda x, edge: torch.where(x >= edge, 1.0, 0.0), + lambda x, edge: 1.0 if x >= edge else 0.0) + + # Three-argument functions + def visitClampFunc(self, ctx): + val = self.visit(ctx.expr(0)) + min_v = self.visit(ctx.expr(1)) + max_v = self.visit(ctx.expr(2)) + # Handle mixed types manually or promote? + if any(self._is_tensor(x) for x in [val, min_v, max_v]): + return torch.clamp(self._promote_to_tensor(val), self._promote_to_tensor(min_v), self._promote_to_tensor(max_v)) + if self._is_list(val): + return [max(min(x, max_v), min_v) for x in val] #TODO: what if min_v or max_v is list + return max(min(val, max_v), min_v) + + def visitLerpFunc(self, ctx): + a = self.visit(ctx.expr(0)) + b = self.visit(ctx.expr(1)) + w = self.visit(ctx.expr(2)) + if any(self._is_tensor(x) for x in [a, b, w]): + # Lerp: a + w*(b-a) + return torch.lerp(self._promote_to_tensor(a), self._promote_to_tensor(b), self._promote_to_tensor(w)) + return a + w * (b - a) + + def visitSmoothstepFunc(self, ctx): + x = self.visit(ctx.expr(0)) + edge0 = self.visit(ctx.expr(1)) + edge1 = self.visit(ctx.expr(2)) + + # t = clamp((x - edge0) / (edge1 - edge0), 0.0, 1.0) + # return t * t * (3.0 - 2.0 * t) + + if any(self._is_tensor(v) for v in [x, edge0, edge1]): + x_t, e0_t, e1_t = self._promote_to_tensor(x), self._promote_to_tensor(edge0), self._promote_to_tensor(edge1) + t = torch.clamp((x_t - e0_t) / (e1_t - e0_t), 0.0, 1.0) + return t * t * (3.0 - 2.0 * t) + + t = max(0.0, min(1.0, (x - edge0) / (edge1 - edge0))) + return t * t * (3.0 - 2.0 * t) + + # Helpers for visiting generic exprs + 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 visitAtomExp(self, ctx): return self.visitChildren(ctx) + def visitExpr(self, ctx): return self.visitChildren(ctx) + + # Original TensorEvalVisitor complex methods (Conv, Map) + # Map, Conv, etc need to handle lists specially now (convert to tensor if expected?) + # or leverage list broadcasting if it makes sense (Conv on a list of images?) + + # For MVP of unification, let's include basic ops and structure, + # and port the complex ones (Conv) carefully. + + # Let's port specific requested functions to verify test suite first. + def visitSMinFunc(self, ctx): + vals = [self.visit(e) for e in ctx.expr()] + + if all(not self._is_tensor(x) and not self._is_list(x) for x in vals): + return min(vals) + + promoted = [self._promote_to_tensor(x) for x in vals] + if len(promoted) == 1: return torch.min(promoted[0]) + return torch.min(torch.stack(torch.broadcast_tensors(*promoted))) + + def visitSMaxFunc(self, ctx): + args = [self.visit(e) for e in ctx.expr()] + if len(args) == 1: + # Check if list or scalar + if not self._is_tensor(args[0]) and not self._is_list(args[0]): return args[0] + if self._is_list(args[0]): return max(args[0]) # max of list + return torch.max(args[0]) # Global max of single tensor + + # Multiple args + if all(not self._is_tensor(x) and not self._is_list(x) for x in args): + return max(args) + + promoted = [self._promote_to_tensor(x) for x in args] + if len(promoted) == 1: return torch.max(promoted[0]) + return torch.max(torch.stack(torch.broadcast_tensors(*promoted))) + + # ========================================== + # Complex Tensor Operations + # ========================================== + + def _fold_nd(self, tsr, spatial_dims): + original_shape = tsr.shape + added_dims = 0 + target_rank = spatial_dims + 2 + while tsr.ndim < target_rank: + tsr = tsr.unsqueeze(1) + added_dims += 1 + 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 + 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:]) + 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._promote_to_tensor(self.visit(ctx.expr(0))) + 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._promote_to_tensor(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 visitPrintShapeFunc(self, ctx): + tsr = self.visit(ctx.expr()) + if hasattr(tsr, 'shape'): print(tsr.shape) + else: print(f"Scalar/List: {tsr}") + return tsr + + def visitSfftFunc(self, ctx): + old_vars = self.variables + self.variables = self.spatial_variables.copy() + try: + val = self._promote_to_tensor(self.visit(ctx.expr())) + dims = tuple(range(val.ndim)) + return torch.fft.fftn(val, dim=dims) + finally: + self.variables = old_vars + + def visitSifftFunc(self, ctx): + old_vars = self.variables + self.variables = self.variables.copy() + device = self.device + + shape_to_use = self.shape if self.shape else (1,1,1,1) + + ndim = len(shape_to_use) + dim_names = ['x', 'y', 'z', 'w', 'v', 'u'] + + k_components = [] + for i in range(ndim): + dim_idx = ndim - 1 - i + size_d = shape_to_use[dim_idx] + 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(*shape_to_use) + + 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) + + self.variables[f'K_dim{dim_idx}'] = values + self.variables[f'F_dim{dim_idx}'] = float(size_d) + k_components.append(values) + + k_sq_sum = torch.zeros(shape_to_use, 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'] + + if 'Kx' in self.variables: + self.variables['frequency_count'] = self.variables.get('Fx', 1.0) + + self.variables = self.variables | generate_dim_variables(k_sq_sum) + + try: + val = self._promote_to_tensor(self.visit(ctx.expr())) + dims = tuple(range(val.ndim)) + return torch.fft.ifftn(val, dim=dims).real + finally: + self.variables = old_vars + + def visitSwapFunc(self, ctx): + tsr = self._promote_to_tensor(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()) if isinstance(dim_t, torch.Tensor) else int(dim_t) + i = int(idx1_t.flatten()[0].item()) if isinstance(idx1_t, torch.Tensor) else int(idx1_t) + j = int(idx2_t.flatten()[0].item()) if isinstance(idx2_t, torch.Tensor) else int(idx2_t) + + 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 _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._promote_to_tensor(self.visit(ctx.expr(0))) + coords = [self._promote_to_tensor(self.visit(ctx.expr(i))) for i in range(1, len(ctx.expr()))] + num_coords = len(coords) + + if num_coords == 0: return tensor + if num_coords > 3: raise ValueError("map() supports max 3 mapping functions.") + + spatial_in_shape = tensor.shape[-num_coords:] + leading_shape = tensor.shape[:-num_coords] + + batch_size = 1 + for s in leading_shape: batch_size *= s + + input_view = tensor.reshape(batch_size, 1, *spatial_in_shape) + norm_coords_list = [] + for i in range(num_coords): + dim_size = spatial_in_shape[i] + norm = self._normalize_coord(coords[i], dim_size) + norm_coords_list.append(norm) + + grid = torch.stack(norm_coords_list[::-1], dim=-1) + grid_spatial_shape = grid.shape[:-1] + + try: + grid_view = grid.reshape(batch_size, *grid_spatial_shape[-(num_coords):], num_coords) + except RuntimeError: + grid_view = grid.expand(batch_size, *([-1] * len(grid_spatial_shape)), -1) + grid_view = grid_view.reshape(batch_size, *grid_view.shape[-(num_coords+1):-1], num_coords) + + if num_coords == 1: + input_final = input_view.reshape(batch_size, 1, 1, -1) + # grid_view is [batch_size, (spatial), 1] + # for 1D it might be just [batch_size, 1] if input was scalar + # we need [batch_size, H_out, W_out, 2] for 2D grid_sample + gv = grid_view + while gv.ndim < 3: gv = gv.unsqueeze(1) # [B, 1, 1] + y_zeros = torch.zeros_like(gv[..., :1]) + grid_final = torch.cat([gv, y_zeros], dim=-1).unsqueeze(1) # [B, 1, 1, 2] + output = F.grid_sample(input_final, grid_final, align_corners=True) + elif num_coords == 2: + grid_final = grid_view.reshape(batch_size, *grid_view.shape[-3:-1], 2) + output = F.grid_sample(input_view, grid_final, align_corners=True) + else: + grid_final = grid_view.reshape(batch_size, *grid_view.shape[-4:-1], 3) + output = F.grid_sample(input_view, grid_final, align_corners=True) + actual_spatial = grid_view.shape[1:-1] + final_shape = list(leading_shape) + list(actual_spatial) + return output.reshape(final_shape) + + def _apply_conv_internal(self, conv_input, kernel_val, kernel_sizes, spatial_dims_count): + """ + Executes the convolution with asymmetric padding support for even kernels. + Input: [Batch, Channel, Spatial...] + Kernel: [Spatial...] (to be promoted/repeated) + """ + in_channels = conv_input.size(1) + + # Calculate Asymmetric Padding for "Same" padding + # Total padding needed = kernel_size - 1 + # Left/Top/Front = (K-1)//2, Right/Bottom/Back = (K-1) - Left + pads = [] + for k in kernel_sizes[::-1]: # F.pad uses reverse order (W, H, D) + p_total = k - 1 + p_low = p_total // 2 + p_high = p_total - p_low + pads.extend([p_low, p_high]) + + # Apply padding + padded_input = torch.nn.functional.pad(conv_input, tuple(pads), mode='constant', value=0) + + # Prepare Kernel + if kernel_val.numel() == 1: + kernel_val = kernel_val.expand(tuple(kernel_sizes)) + elif kernel_val.ndim != spatial_dims_count: + kernel_val = kernel_val.reshape(tuple(kernel_sizes)) + + final_kernel = kernel_val.unsqueeze(0).unsqueeze(0) + final_kernel = final_kernel.to(conv_input.dtype) + # Repeat for Depthwise-like behavior: Weight [OutC, InC/Groups, K...] + # result = conv(Groups=InC, InC=InC) -> Weight [InC, 1, K...] + final_kernel = final_kernel.repeat(in_channels, 1, *([1]*spatial_dims_count)) + + conv_fn = F.conv1d if spatial_dims_count == 1 else (F.conv2d if spatial_dims_count == 2 else F.conv3d) + + # padding=0 because we padded explicitly via F.pad + result = conv_fn(padded_input, final_kernel, padding=0, groups=in_channels) + return result + + def visitConvFunc(self, ctx): + tensor = self._promote_to_tensor(self.visit(ctx.expr(0))) + num_args = len(ctx.expr()) + if num_args < 3: raise ValueError("conv() requires at least 3 arguments") + + kernel_arg_idx = num_args - 1 + + spatial_dims_count = num_args - 2 + if spatial_dims_count not in [1, 2, 3]: + raise ValueError(f"conv() supports 1D, 2D, or 3D. Found {spatial_dims_count}") + + kernel_sizes = [] + for i in range(1, 1 + spatial_dims_count): + val = self.visit(ctx.expr(i)) + if isinstance(val, torch.Tensor): val = int(val.flatten()[0].item()) + kernel_sizes.append(val) + + # Prepare context for kernel + coords = [torch.arange(s, device=self.device).float() - (s//2) for s in kernel_sizes] + grid = torch.meshgrid(*coords, indexing='ij') + + dim_names = ['x','y','z'] + old_vars = self.variables + self.variables = self.variables.copy() + for i in range(spatial_dims_count): + self.variables[f'k{dim_names[i].lower()}'] = grid[i] + self.variables[f'k{dim_names[i].upper()}'] = grid[i] + + original_shape = self.shape + self.shape = tuple(kernel_sizes) + + try: + kernel_val = self._promote_to_tensor(self.visit(ctx.expr(kernel_arg_idx))) + finally: + self.shape = original_shape + self.variables = old_vars + + # --- Dimension Management (Standardizing to BHWC internally) --- + + # 1. Detect Layout (Channels-First vs Channels-Last) + is_channels_first = False + if tensor.ndim >= 3 and tensor.shape[-1] > 4: + is_channels_first = True + + # 2. Convert Channels-First [B, C, S...] to Channels-Last [B, S..., C] + if is_channels_first: + if spatial_dims_count == 1 and tensor.ndim == 3: + tensor = tensor.permute(0, 2, 1) + elif spatial_dims_count == 2 and tensor.ndim == 4: + tensor = tensor.permute(0, 2, 3, 1) + elif spatial_dims_count == 3: + if tensor.ndim == 4: tensor = tensor.unsqueeze(-1) # [B, C, H, W] -> [B, C, H, W, 1] (C is Depth) + elif tensor.ndim == 5: tensor = tensor.permute(0, 2, 3, 4, 1) + + # 3. Handle user rule: "last 4 dimensions as d,v,h,c" for ndim=4 + # At this point, for 3D conv on Latents [B, D, H, W, 1], we have 5 dims. + # Ensure we have N+2 dimensions for conv logic + if tensor.ndim == spatial_dims_count + 1: + tensor = tensor.unsqueeze(-1) # Add C=1 + elif tensor.ndim == spatial_dims_count: + tensor = tensor.unsqueeze(0).unsqueeze(-1) # Add B=1, C=1 + + # 4. Partition Dimensions + in_channels = tensor.size(-1) + batch_end_idx = tensor.ndim - 1 - spatial_dims_count + batch_shape = tensor.shape[:batch_end_idx] + spatial_shape = tensor.shape[batch_end_idx:-1] + channels_shape = (tensor.shape[-1],) + + # 5. Flatten / Permute to [N, C, S...] for helper + total_batch = 1 + for s in batch_shape: total_batch *= s + flat_input = tensor.reshape(total_batch, *spatial_shape, in_channels) + + permute_order = [0, spatial_dims_count + 1] + list(range(1, spatial_dims_count + 1)) + conv_input = flat_input.permute(*permute_order) + + # 6. Call pure kernel runner (with padding fix) + result = self._apply_conv_internal(conv_input, kernel_val, kernel_sizes, spatial_dims_count) + + # 7. Reverse Permute / Un-flatten / Un-pad + # Helper returned [N, C, S...] + result_permute = [0] + list(range(2, 2+spatial_dims_count)) + [1] + out_flat = result.permute(*result_permute) + + final_shape = batch_shape + spatial_shape + channels_shape + out = out_flat.reshape(final_shape) + + # 8. Restore Channels-First if needed + if is_channels_first: + if spatial_dims_count == 1 and out.ndim == 3: + out = out.permute(0, 2, 1) + elif spatial_dims_count == 2 and out.ndim == 4: + out = out.permute(0, 3, 1, 2) + elif spatial_dims_count == 3: + if out.ndim == 5 and out.shape[-1] == 1: + out = out.squeeze(-1) + elif out.ndim == 5: + out = out.permute(0, 4, 1, 2, 3) + + return out + + # Copied from TensorEvalVisitor but using unified logic where applicable + # Note: For Conv/Map, we stick to Tensor logic mostly, but if args are lists we might error or auto-stack. + # The user mentioned: "easy ability to use it [list] in conv after reshaping". + # This implies conv(list, ...) might be useful. + # But usually conv input is a tensor. + # If list is passed to conv(A ...), A must be tensor? + # Or conv([img1, img2], ...) -> [conv(img1), conv(img2)]? + # Broadcasting logic handles list inputs naturally if we map `visit` over list. + # But `conv` is a custom Visitor method, not routed via `_bin_op`. + # We would need to implement list handling inside `visitConvFunc`. + + # Implementing generic fallback for missing methods to avoid crashes during dev? + # No, better fail. diff --git a/more_math/helper_functions.py b/more_math/helper_functions.py index b019bca..422f7d2 100644 --- a/more_math/helper_functions.py +++ b/more_math/helper_functions.py @@ -25,41 +25,32 @@ def parse_expr(expr: str): def eval_tensor_expr(expr: str, variables: dict, shape: tuple, device=None): - """Parse and evaluate a tensor math expression. - - Args: - expr: Math expression string - variables: Dict of variable names to tensor/scalar values - shape: Shape tuple for the TensorEvalVisitor - device: Optional device override - - Returns: - Result tensor from evaluating the expression - """ - from .Parser.TensorEvalVisitor import TensorEvalVisitor + """Parse and evaluate a tensor math expression.""" + from .Parser.UnifiedMathVisitor import UnifiedMathVisitor tree = parse_expr(expr) - visitor = TensorEvalVisitor(variables, shape, device=device) + visitor = UnifiedMathVisitor(variables, shape, device=device) return visitor.visit(tree) def eval_tensor_expr_with_tree(tree, variables: dict, shape: tuple, device=None): - """Evaluate a pre-parsed expression tree with TensorEvalVisitor.""" - from .Parser.TensorEvalVisitor import TensorEvalVisitor - visitor = TensorEvalVisitor(variables, shape, device=device) + """Evaluate a pre-parsed expression tree with UnifiedMathVisitor.""" + from .Parser.UnifiedMathVisitor import UnifiedMathVisitor + visitor = UnifiedMathVisitor(variables, shape, device=device) return visitor.visit(tree) def eval_float_expr(expr: str, variables: dict): """Parse and evaluate a float math expression.""" - from .Parser.FloatEvalVisitor import FloatEvalVisitor + from .Parser.UnifiedMathVisitor import UnifiedMathVisitor tree = parse_expr(expr) - visitor = FloatEvalVisitor(variables) + # Float eval context often has no shape. Pass None/Empty. + visitor = UnifiedMathVisitor(variables, shape=None) return visitor.visit(tree) def eval_float_expr_with_tree(tree, variables: dict): - """Evaluate a pre-parsed expression tree with FloatEvalVisitor.""" - from .Parser.FloatEvalVisitor import FloatEvalVisitor - visitor = FloatEvalVisitor(variables) + """Evaluate a pre-parsed expression tree with UnifiedMathVisitor.""" + from .Parser.UnifiedMathVisitor import UnifiedMathVisitor + visitor = UnifiedMathVisitor(variables, shape=None) return visitor.visit(tree) @@ -86,7 +77,7 @@ def comonLazy(expr, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0): need_eval.append(token.text) return need_eval -def generate_dim_variables(tensor): +def generate_dim_variables(tensor: torch.Tensor): """Generate index and size tensors for each dimension of the input tensor.""" variables = {} for dim, size in enumerate(tensor.shape):