"""A scalar reverse-mode autodiff engine designed for inspection and tests.""" from __future__ import annotations import math from typing import Callable, Iterable class Value: def __init__( self, data: float, parents: Iterable["Value"] = (), operation: str = "leaf", ): self.data = float(data) self.grad = 0.0 self.parents = tuple(parents) self.operation = operation self._backward: Callable[[], None] = lambda: None @staticmethod def wrap(value: float | "Value") -> "Value": return value if isinstance(value, Value) else Value(value) def __add__(self, other: float | "Value") -> "Value": other = self.wrap(other) output = Value(self.data + other.data, (self, other), "add") def backward() -> None: self.grad += output.grad other.grad += output.grad output._backward = backward return output def __radd__(self, other: float | "Value") -> "Value": return self + other def __mul__(self, other: float | "Value") -> "Value": other = self.wrap(other) output = Value(self.data * other.data, (self, other), "multiply") def backward() -> None: self.grad += other.data * output.grad other.grad += self.data * output.grad output._backward = backward return output def __rmul__(self, other: float | "Value") -> "Value": return self * other def tanh(self) -> "Value": value = math.tanh(self.data) output = Value(value, (self,), "tanh") def backward() -> None: self.grad += (1 - value**2) * output.grad output._backward = backward return output def backward(self) -> None: order: list[Value] = [] seen: set[int] = set() def visit(node: Value) -> None: if id(node) in seen: return seen.add(id(node)) for parent in node.parents: visit(parent) order.append(node) visit(self) self.grad = 1.0 for node in reversed(order): node._backward() def central_difference(function: Callable[[float], float], x: float, epsilon=1e-5): if epsilon <= 0: raise ValueError("epsilon must be positive") return (function(x + epsilon) - function(x - epsilon)) / (2 * epsilon)