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