"""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)