import ast import re import numpy as np from number_format import format_num as _scalar_format_num def format_num(val, sig_figs=4): if not np.isfinite(val): return str(float(val)) return _scalar_format_num(val, sig_figs=sig_figs) MAX_ARRAY_RANK = 2 MAX_TOTAL_ELEMENTS = 50000 MAX_PREVIEW_ELEMENTS = 4 MAX_EXPR_LEN = 200 MAX_AST_NODES = 80 MAX_AST_DEPTH = 15 MAX_NUMERIC_MAGNITUDE = 1e9 MAX_EXPONENT_MAGNITUDE = 50 TENSOR_ALLOWED_FUNCS = { "sigmoid", "relu", "tanh", "exp", "log", "sum", "mean", "dot", "max", "min", "softmax", "pick", "nograd", } class TensorASTValidator(ast.NodeVisitor): def __init__(self): self.node_count = 0 self.max_depth = 0 self.current_depth = 0 def visit(self, node): self.node_count += 1 if self.node_count > MAX_AST_NODES: raise ValueError(f"Expression exceeds complexity limit ({MAX_AST_NODES} nodes).") self.current_depth += 1 if self.current_depth > self.max_depth: self.max_depth = self.current_depth if self.max_depth > MAX_AST_DEPTH: raise ValueError(f"Expression nesting too deep (max depth {MAX_AST_DEPTH}).") res = super().visit(node) self.current_depth -= 1 return res def generic_visit(self, node): allowed_types = ( ast.Expression, ast.BinOp, ast.UnaryOp, ast.Add, ast.Sub, ast.Mult, ast.Div, ast.Pow, ast.MatMult, ast.UAdd, ast.USub, ast.Name, ast.Constant, ast.Call, ast.Load, ast.List, ast.Subscript, ) if hasattr(ast, "Num"): allowed_types = allowed_types + (ast.Num,) if not isinstance(node, allowed_types): raise ValueError(f"Disallowed construct: {type(node).__name__}") super().generic_visit(node) def visit_Call(self, node): if node.keywords: raise ValueError("Keyword arguments are not allowed in function calls.") self.generic_visit(node) from tensor_formatter import format_tensor_value, format_tensor_shape, DEFAULT_TENSOR_ELEMENT_BUDGET def format_tensor(val, sig_figs=4): """Canonical tensor formatting following the 16-element budget.""" if val is None: return "" return format_tensor_value(val, element_budget=DEFAULT_TENSOR_ELEMENT_BUDGET, include_shape_if_truncated=True) def format_tensor_badge(val, sig_figs=4): """Badge formatting for edge labels: represents value canonically without extra shape suffix.""" if val is None: return "" # Edge labels show value without trailing shape metadata to stay compact return format_tensor_value(val, element_budget=DEFAULT_TENSOR_ELEMENT_BUDGET, include_shape_if_truncated=False) class TensorNode: def __init__(self, uid, name, op_type, is_input=False, is_constant=False, const_val=None): self.uid = uid self.name = name self.op_type = op_type # 'var', 'const', 'add', 'sub', 'mul', 'div', 'pow', 'matmul', 'func' self.is_input = is_input self.is_constant = is_constant self.const_val = const_val self.args = [] # list of child TensorNode self.val = None # float or np.ndarray self.grad = None # float or np.ndarray (accumulated adjoint) self.shape = None self.func_name = None # for func op self.pow_exp = None # for pow op self.pick_index = None # nondifferentiable literal metadata self.live = False # reached by a non-stopped backward path # Backward contributions to operands: list of dicts # Each record carries an operand occurrence, a typed rule, and its value. self.backward_rules = [] class TensorComputationGraph: def __init__(self): self.nodes = [] self.name_to_input = {} self.op_counter = 0 self.inp_counter = 0 self.cst_counter = 0 self.orig_expr_str = "" self.root = None def _get_op_display_name(self): number = self.op_counter while f"n{number}" in self.used_names: number += 1 name = f"n{number}" self.used_names.add(name) return name def build_from_ast(self, expr_str): if not isinstance(expr_str, str): raise ValueError("Expression must be a string.") cleaned = expr_str.strip() if not cleaned: raise ValueError("Expression cannot be empty.") if len(cleaned) > MAX_EXPR_LEN: raise ValueError(f"Expression too long ({len(cleaned)} chars, max {MAX_EXPR_LEN}).") try: parsed = ast.parse(cleaned, mode="eval") except SyntaxError as e: raise ValueError(f"Syntax error in expression: {e.msg}") validator = TensorASTValidator() validator.visit(parsed) self.used_names = {n.id for n in ast.walk(parsed) if isinstance(n, ast.Name)} self.input_names = self.used_names.copy() self.orig_expr_str = cleaned self.nodes = [] self.name_to_input = {} self.op_counter = 0 self.inp_counter = 0 self.cst_counter = 0 self.root = self._build_recursive(parsed.body) if not self.root.is_input: root_name = "out" while root_name in self.input_names: root_name += "_result" self.root.name = root_name return self def _extract_constant_scalar(self, node): """ Helper to extract signed numeric scalar literal from an AST node (e.g. 2, -2, +0.5). Returns float value or None if node is not a constant scalar. """ if isinstance(node, ast.Constant): if isinstance(node.value, bool) or not isinstance(node.value, (int, float)): return None return float(node.value) if hasattr(ast, "Num") and isinstance(node, ast.Num): return float(node.n) if isinstance(node, ast.UnaryOp): val = self._extract_constant_scalar(node.operand) if val is None: return None if isinstance(node.op, ast.UAdd): return +val elif isinstance(node.op, ast.USub): return -val return None def _build_selection(self, value, index): sign = 1 if isinstance(index, ast.UnaryOp) and isinstance(index.op, (ast.USub, ast.UAdd)): sign = -1 if isinstance(index.op, ast.USub) else 1 index = index.operand if not isinstance(index, ast.Constant) or type(index.value) is not int: raise ValueError("Vector index must be an integer literal (zero-based; negatives count from the end).") child = self._build_recursive(value) self.op_counter += 1 result = TensorNode(f"op_{self.op_counter}", self._get_op_display_name(), "func") result.func_name = "pick" result.pick_index = sign * index.value result.args = [child] self.nodes.append(result) return result def _build_recursive(self, node): if isinstance(node, ast.Subscript): return self._build_selection(node.value, node.slice) if isinstance(node, ast.Constant): if isinstance(node.value, bool): raise ValueError("Boolean constants are not allowed.") if not isinstance(node.value, (int, float)): raise ValueError(f"Unsupported constant type: {type(node.value).__name__}") val = float(node.value) if abs(val) > MAX_NUMERIC_MAGNITUDE: raise ValueError(f"Numeric constant {val} exceeds maximum allowed magnitude ({MAX_NUMERIC_MAGNITUDE}).") self.cst_counter += 1 c_node = TensorNode( uid=f"cst_{self.cst_counter}_{str(val).replace('-', 'neg_').replace('.', '_')}", name=format_num(val), op_type="const", is_input=True, is_constant=True, const_val=val, ) self.nodes.append(c_node) return c_node if hasattr(ast, "Num") and isinstance(node, ast.Num): val = float(node.n) if abs(val) > MAX_NUMERIC_MAGNITUDE: raise ValueError(f"Numeric constant {val} exceeds maximum allowed magnitude.") self.cst_counter += 1 c_node = TensorNode( uid=f"cst_{self.cst_counter}_{str(val).replace('-', 'neg_').replace('.', '_')}", name=format_num(val), op_type="const", is_input=True, is_constant=True, const_val=val, ) self.nodes.append(c_node) return c_node if isinstance(node, ast.List): from input_parser import _validate_and_convert_literal arr_val = np.array(_validate_and_convert_literal(node), dtype=np.float64) self.cst_counter += 1 c_node = TensorNode( uid=f"cst_{self.cst_counter}_arr", name=format_tensor(arr_val), op_type="const", is_input=True, is_constant=True, const_val=arr_val, ) self.nodes.append(c_node) return c_node if isinstance(node, ast.Name): var_name = node.id if not re.match(r"^[a-zA-Z][a-zA-Z0-9_]*$", var_name): raise ValueError(f"Invalid variable identifier: {var_name}") if var_name in self.name_to_input: return self.name_to_input[var_name] self.inp_counter += 1 inp_node = TensorNode( uid=f"inp_{self.inp_counter}_{var_name}", name=var_name, op_type="var", is_input=True, is_constant=False, ) self.name_to_input[var_name] = inp_node self.nodes.append(inp_node) return inp_node if isinstance(node, ast.UnaryOp): if isinstance(node.op, ast.UAdd): return self._build_recursive(node.operand) elif isinstance(node.op, ast.USub): # Desugar -x as 0 - x or negate op child = self._build_recursive(node.operand) self.op_counter += 1 op_node = TensorNode( uid=f"op_{self.op_counter}", name=self._get_op_display_name(), op_type="neg", is_input=False, ) op_node.args = [child] self.nodes.append(op_node) return op_node raise ValueError(f"Unsupported unary operator: {type(node.op).__name__}") if isinstance(node, ast.BinOp): if isinstance(node.op, ast.Pow): left_node = self._build_recursive(node.left) # Check for signed numeric constant exponent exp_val = self._extract_constant_scalar(node.right) if exp_val is None: raise ValueError("Tensor power currently requires a finite scalar constant exponent.") if abs(exp_val) > MAX_EXPONENT_MAGNITUDE: raise ValueError(f"Exponent magnitude {exp_val} exceeds limit of {MAX_EXPONENT_MAGNITUDE}.") self.cst_counter += 1 right_node = TensorNode( uid=f"cst_{self.cst_counter}_{str(exp_val).replace('-', 'neg_').replace('.', '_')}", name=format_num(exp_val), op_type="const", is_input=True, is_constant=True, const_val=exp_val, ) self.nodes.append(right_node) self.op_counter += 1 op_name = self._get_op_display_name() op_node = TensorNode( uid=f"op_{self.op_counter}", name=op_name, op_type="pow", is_input=False, ) op_node.args = [left_node, right_node] op_node.pow_exp = float(exp_val) self.nodes.append(op_node) return op_node left_node = self._build_recursive(node.left) right_node = self._build_recursive(node.right) self.op_counter += 1 op_name = self._get_op_display_name() if isinstance(node.op, ast.Add): op_type = "add" elif isinstance(node.op, ast.Sub): op_type = "sub" elif isinstance(node.op, ast.Mult): op_type = "mul" elif isinstance(node.op, ast.Div): op_type = "div" elif isinstance(node.op, ast.MatMult): op_type = "matmul" else: raise ValueError(f"Unsupported binary operator: {type(node.op).__name__}") op_node = TensorNode( uid=f"op_{self.op_counter}", name=op_name, op_type=op_type, is_input=False, ) op_node.args = [left_node, right_node] self.nodes.append(op_node) return op_node if isinstance(node, ast.Call): if not isinstance(node.func, ast.Name): raise ValueError("Only direct function calls (e.g. sum(x), dot(u, v)) are allowed.") if node.keywords: raise ValueError("Keyword arguments are not allowed in function calls.") func_name = node.func.id if func_name not in TENSOR_ALLOWED_FUNCS: raise ValueError(f"Function '{func_name}' is not supported in tensor expressions: {sorted(TENSOR_ALLOWED_FUNCS)}") if func_name == "pick": if len(node.args) != 2: raise ValueError("Function 'pick' requires exactly 2 arguments.") return self._build_selection(node.args[0], node.args[1]) elif func_name == "dot": if len(node.args) != 2: raise ValueError(f"Function 'dot' requires exactly 2 arguments, got {len(node.args)}.") left_node = self._build_recursive(node.args[0]) right_node = self._build_recursive(node.args[1]) self.op_counter += 1 op_node = TensorNode( uid=f"op_{self.op_counter}", name=self._get_op_display_name(), op_type="dot", is_input=False, ) op_node.args = [left_node, right_node] self.nodes.append(op_node) return op_node else: if len(node.args) != 1: raise ValueError(f"Function '{func_name}' requires exactly 1 argument, got {len(node.args)}.") child_node = self._build_recursive(node.args[0]) self.op_counter += 1 op_node = TensorNode( uid=f"op_{self.op_counter}", name=self._get_op_display_name(), op_type="func", is_input=False, ) op_node.func_name = func_name op_node.args = [child_node] self.nodes.append(op_node) return op_node raise ValueError(f"Unsupported syntax: {type(node).__name__}") def evaluate(self, feed_dict, forward_only=False): # 1. Forward pass for node in self.nodes: if node.is_input: if node.is_constant: node.val = node.const_val elif node.name in feed_dict: node.val = feed_dict[node.name] else: raise ValueError(f"Missing numerical value for variable '{node.name}'") # Validate rank & shape if isinstance(node.val, np.ndarray): if node.val.ndim > MAX_ARRAY_RANK: raise ValueError(f"Variable '{node.name}' rank {node.val.ndim} exceeds maximum {MAX_ARRAY_RANK}.") if node.val.size > MAX_TOTAL_ELEMENTS: raise ValueError(f"Variable '{node.name}' size ({node.val.size}) exceeds limit ({MAX_TOTAL_ELEMENTS}).") node.shape = node.val.shape else: node.val = np.float64(node.val) node.shape = () else: # Evaluate intermediate operation self._evaluate_forward_op(node) # Enforce scalar root restriction in both modes if self.root.shape != (): raise ValueError("The final expression must produce a scalar. Use sum(...) or mean(...), or define a scalar loss.") if forward_only: return # 2. Backward pass (VJP Adjoint Propagation) for node in self.nodes: if node.shape == (): node.grad = 0.0 else: node.grad = np.zeros(node.shape, dtype=np.float64) node.backward_rules = [] # Seed output adjoint with 1.0 (since root is scalar ()) self.root.grad = 1.0 # A node is "live" if some path of real (non-stopped) gradient flow # reaches it from the root. nograd() blocks its operand from becoming # live, so backward computation never recurses past it. for node in self.nodes: node.live = False self.root.live = True for node in reversed(self.nodes): if node.is_input or not node.live: continue is_nograd = node.op_type == "func" and node.func_name == "nograd" if is_nograd: continue # hard stop: no backward step, no adjoint edge self._evaluate_backward_op(node) for child in node.args: child.live = True @np.errstate(all="ignore") def _evaluate_forward_op(self, node): op = node.op_type if op == "neg": child = node.args[0] node.val = -child.val node.shape = child.shape elif op == "add": a, b = node.args[0], node.args[1] self._validate_elementwise_shapes(a, b, "+") node.val = a.val + b.val node.shape = node.val.shape if isinstance(node.val, np.ndarray) else () elif op == "sub": a, b = node.args[0], node.args[1] self._validate_elementwise_shapes(a, b, "-") node.val = a.val - b.val node.shape = node.val.shape if isinstance(node.val, np.ndarray) else () elif op == "mul": a, b = node.args[0], node.args[1] self._validate_elementwise_shapes(a, b, "*") node.val = a.val * b.val node.shape = node.val.shape if isinstance(node.val, np.ndarray) else () elif op == "div": a, b = node.args[0], node.args[1] self._validate_elementwise_shapes(a, b, "/") node.val = np.divide(a.val, b.val) node.shape = node.val.shape if isinstance(node.val, np.ndarray) else () elif op == "pow": base = node.args[0] exp = node.pow_exp node.val = np.power(base.val, exp) node.shape = node.val.shape if isinstance(node.val, np.ndarray) else () elif op == "matmul": a, b = node.args[0], node.args[1] # Reject scalar operands for @ if a.shape == () or b.shape == (): raise ValueError("Matrix multiplication (@) requires vector or matrix operands, not scalars.") # Vector-Vector -> Scalar if a.val.ndim == 1 and b.val.ndim == 1: if a.val.shape[0] != b.val.shape[0]: raise ValueError(f"Vector dot product dimension mismatch: {a.val.shape[0]} vs {b.val.shape[0]}") node.val = float(np.dot(a.val, b.val)) node.shape = () # Matrix-Vector -> Vector elif a.val.ndim == 2 and b.val.ndim == 1: if a.val.shape[1] != b.val.shape[0]: raise ValueError(f"Matrix-vector dimension mismatch: {a.val.shape} @ {b.val.shape}") node.val = np.matmul(a.val, b.val) node.shape = node.val.shape # Vector-Matrix -> Vector elif a.val.ndim == 1 and b.val.ndim == 2: if a.val.shape[0] != b.val.shape[0]: raise ValueError(f"Vector-matrix dimension mismatch: {a.val.shape} @ {b.val.shape}") node.val = np.matmul(a.val, b.val) node.shape = node.val.shape # Matrix-Matrix -> Matrix elif a.val.ndim == 2 and b.val.ndim == 2: if a.val.shape[1] != b.val.shape[0]: raise ValueError(f"Matrix-matrix dimension mismatch: {a.val.shape} @ {b.val.shape}") node.val = np.matmul(a.val, b.val) node.shape = node.val.shape else: raise ValueError(f"Unsupported matmul ranks: {a.val.ndim}, {b.val.ndim}") elif op == "dot": u, v = node.args[0], node.args[1] if u.shape == () or v.shape == () or u.val.ndim != 1 or v.val.ndim != 1: raise ValueError("Function 'dot(u, v)' requires two 1D vectors.") if u.val.shape[0] != v.val.shape[0]: raise ValueError(f"Length mismatch in dot(u, v): {u.val.shape[0]} vs {v.val.shape[0]}") node.val = float(np.dot(u.val, v.val)) node.shape = () elif op == "func": child = node.args[0] f = node.func_name if f == "pick": if len(child.shape) != 1: raise ValueError("pick requires a 1D vector.") if not -child.shape[0] <= node.pick_index < child.shape[0]: raise ValueError(f"pick index {node.pick_index} is out of range for vector length {child.shape[0]}.") node.val = float(child.val[node.pick_index]) node.shape = () elif f == "softmax": if len(child.shape) != 1 or child.shape[0] == 0: raise ValueError("softmax requires a nonempty 1D vector.") weights = np.exp(child.val - np.max(child.val)) node.val = weights / np.sum(weights) node.shape = child.shape elif f == "max": node.val = float(np.max(child.val)) node.shape = () elif f == "min": node.val = float(np.min(child.val)) node.shape = () elif f == "sum": node.val = float(np.sum(child.val)) node.shape = () elif f == "mean": node.val = float(np.mean(child.val)) node.shape = () elif f == "sigmoid": # Numerically stable sigmoid c_val = child.val node.val = np.where(c_val >= 0, 1.0 / (1.0 + np.exp(-c_val)), np.exp(c_val) / (1.0 + np.exp(c_val))) if isinstance(node.val, np.ndarray) and node.val.ndim == 0: node.val = np.float64(node.val) node.shape = child.shape elif f == "relu": node.val = np.maximum(0.0, child.val) if isinstance(node.val, np.ndarray) and node.val.ndim == 0: node.val = np.float64(node.val) node.shape = child.shape elif f == "tanh": node.val = np.tanh(child.val) if isinstance(node.val, np.ndarray) and node.val.ndim == 0: node.val = np.float64(node.val) node.shape = child.shape elif f == "exp": node.val = np.exp(child.val) if isinstance(node.val, np.ndarray) and node.val.ndim == 0: node.val = np.float64(node.val) node.shape = child.shape elif f == "log": node.val = np.log(child.val) if isinstance(node.val, np.ndarray) and node.val.ndim == 0: node.val = np.float64(node.val) node.shape = child.shape elif f == "nograd": node.val = child.val node.shape = child.shape else: raise ValueError(f"Unknown function: {f}") if isinstance(node.val, np.ndarray) and node.val.size > MAX_TOTAL_ELEMENTS: raise ValueError(f"Intermediate array size ({node.val.size}) exceeds limit ({MAX_TOTAL_ELEMENTS}).") def _validate_elementwise_shapes(self, a, b, op_symbol): if a.shape == b.shape: return # Allow scalar-to-array broadcasting if a.shape == () or b.shape == (): return # Reject vector-to-matrix and unequal array broadcasting raise ValueError( f"Unsupported broadcasting for '{op_symbol}': operands have shapes {a.shape} and {b.shape}. " "Only equal-shape operands and scalar-to-array broadcasting are supported." ) @np.errstate(all="ignore") def _evaluate_backward_op(self, node): from backward_rules import ScalarLocalRule, tensor_rule from computation_model import tensor_operation operation = tensor_operation(node) values = {a.uid: a.val for a in node.args} values[node.uid] = node.val values['adjoint:' + node.uid] = node.grad occurrences = range(1) if node.op_type == 'pow' else range(len(node.args)) for occurrence in occurrences: child = node.args[occurrence] rule = tensor_rule(operation, occurrence) contribution = rule.application.evaluate(values) if child.shape == (): contribution = float(contribution) record = {'arg_node': child, 'operand_occurrence': occurrence, 'rule': rule, 'contrib_val': contribution} if isinstance(rule, ScalarLocalRule): record['local_val'] = rule.mapped_derivative.evaluate(values) node.backward_rules.append(record) child.grad = child.grad + contribution def to_computation_model(self, output_name="out", forward_only=False): from computation_model import build_computation_model return build_computation_model(self, "tensor", output_name, forward_only) def to_diagram_model(self, output_name="out", forward_only=False): """Compatibility alias; the store also serves walkthroughs and exporters.""" return self.to_computation_model(output_name, forward_only) def to_mermaid(self, output_name="out", forward_only=False, grad_symbol=None, diff_notation="code"): from diagram_renderer import render_diagram if grad_symbol is not None and grad_symbol != 1: raise ValueError("The scalar output seeds backward computation with adjoint 1.") return render_diagram(self.to_diagram_model(output_name, forward_only), diff_notation)