Spaces:
Running
Running
Download tensor_engine.py from simon-clmtd/autodiff-visualizer: direct link, hf CLI and curl.
- Browser
- Download file 27.6 kB
-
https://huggingface.co/spaces/simon-clmtd/autodiff-visualizer/resolve/main/tensor_engine.py
- Command line
-
hf download hf://spaces/simon-clmtd/autodiff-visualizer/tensor_engine.py
-
curl -L -o tensor_engine.py https://huggingface.co/spaces/simon-clmtd/autodiff-visualizer/resolve/main/tensor_engine.py
27.6 kB
| 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 | |
| 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." | |
| ) | |
| 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) | |