autodiff-visualizer / tensor_engine.py
Simon Clematide
Refactor backward explanations into shared structured model
8565ab7
Raw History Blame Contribute Delete
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
@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)