autodiff-visualizer / math_expression.py
Simon Clematide
Refactor backward explanations into shared structured model
8565ab7
Raw History Blame Contribute Delete
5.43 kB
"""Small renderer-independent expression vocabulary for teaching reverse mode.
References use stable identities; names are presentation metadata. No HTML or
notation strings belong here. This is deliberately not a computer algebra system.
"""
from dataclasses import dataclass, replace
from typing import Any, Mapping
import numpy as np
def snapshot(value):
if isinstance(value, np.ndarray):
value = value.copy()
value.flags.writeable = False
return value
@dataclass(frozen=True, eq=False)
class Expr:
kind: str
args: tuple['Expr', ...] = ()
value: Any = None
key: str = ''
name: str = ''
def substitute(self, bindings: Mapping[str, 'Expr']) -> 'Expr':
if self.kind in {'ref', 'adjoint', 'slot'} and self.key in bindings:
return bindings[self.key]
return replace(self, args=tuple(a.substitute(bindings) for a in self.args))
def bind(self, values: Mapping[str, Any]) -> 'Expr':
return self.substitute({key: literal(value) for key, value in values.items()})
def numerical_application(self, values):
"""Substitute values, evaluating local factors while retaining adjoint flow."""
if self.kind == 'adjoint':
return literal(values[self.key])
if self.kind in {'ref', 'slot'}:
return literal(values[self.key])
if self.kind == 'literal':
return self
if not self.has_adjoint():
return literal(self.evaluate(values))
return replace(self, args=tuple(a.numerical_application(values) for a in self.args))
def has_adjoint(self):
return self.kind == 'adjoint' or any(a.has_adjoint() for a in self.args)
def to_data(self):
"""Portable structure for future exporters (including compact/TikZ views)."""
value = self.value
if isinstance(value, np.ndarray):
value = value.tolist()
elif isinstance(value, np.generic):
value = value.item()
return {'kind': self.kind, 'args': [a.to_data() for a in self.args],
'value': value, 'key': self.key, 'name': self.name}
@np.errstate(all='ignore')
def evaluate(self, values: Mapping[str, Any] | None = None):
if self.kind == 'literal':
return self.value
if self.kind in {'ref', 'adjoint', 'slot'}:
return values[self.key]
args = [a.evaluate(values) for a in self.args]
if self.kind == 'add':
return sum(args)
if self.kind in {'mul', 'elementwise', 'scale'}:
result = args[0]
for arg in args[1:]:
result = result * arg
return result
operations = {
'sub': np.subtract, 'neg': np.negative, 'div': np.divide,
'pow': np.power, 'matmul': np.matmul, 'outer': np.outer,
'dot': np.dot, 'transpose': np.transpose, 'sum': np.sum,
'mean': np.mean, 'max': np.max, 'min': np.min,
'ones_like': np.ones_like, 'zeros_like': np.zeros_like,
'size': np.size, 'equal': lambda a, b: np.asarray(a == b, dtype=float),
'positive': lambda a: np.asarray(a > 0, dtype=float),
'pick': lambda a, i: a[int(i)],
'one_hot': _one_hot,
'sigmoid': lambda a: 1 / (1 + np.exp(-a)),
'relu': lambda a: np.maximum(a, 0),
'softmax': lambda a: np.exp(a - np.max(a)) / np.sum(np.exp(a - np.max(a))),
'nograd': lambda a: a,
'sin': np.sin, 'cos': np.cos, 'tan': np.tan,
'asin': np.arcsin, 'acos': np.arccos, 'atan': np.arctan,
'sinh': np.sinh, 'cosh': np.cosh, 'tanh': np.tanh,
'exp': np.exp, 'log': np.log, 'Abs': np.abs, 'sign': np.sign,
}
if self.kind not in operations:
raise ValueError(f'Expression is not evaluable: {self.kind}')
return operations[self.kind](*args)
def _one_hot(value, index):
result = np.zeros_like(value, dtype=float)
result[int(index)] = 1
return result
def literal(value):
return Expr('literal', value=snapshot(value))
def ref(quantity):
return Expr('ref', key=quantity.uid, name=quantity.name)
def adjoint(quantity):
return Expr('adjoint', key='adjoint:' + quantity.uid, name=quantity.name)
def node(kind, *args):
return Expr(kind, tuple(args))
def from_sympy(expr, symbols):
"""Convert the scalar engine's derivative without flattening it into text."""
if expr in symbols:
return symbols[expr]
if expr.is_Number:
return literal(float(expr))
if expr.is_Symbol:
raise ValueError(f'Unbound scalar symbol: {expr}')
kind = {'Add': 'add', 'Mul': 'mul', 'Pow': 'pow'}.get(expr.func.__name__, expr.func.__name__)
return node(kind, *(from_sympy(a, symbols) for a in expr.args))
def operation_expression(operation, operands=None):
args = tuple(ref(q) for q in operation.operands) if operands is None else tuple(operands)
kind = operation.kind
if kind == 'func':
kind = operation.function
if kind == 'mul' and any(q.shape for q in operation.operands):
kind = 'elementwise' if all(q.shape for q in operation.operands) else 'scale'
if kind == 'pick':
return node('pick', args[0], literal(operation.index))
if kind == 'pow' and len(args) == 1:
args += (literal(operation.exponent),)
return node(kind, *args)