Concept 02: Building an Autograd Engine in Pure Java
In Concept 25, we traced gradients through a computational graph by hand. Now, we will write a complete Automatic Differentiation Engine (Autograd)—the core foundation of PyTorch—in about 30 lines of pure Python!
Open the interactive demo below to build and evaluate custom computational expressions and watch the autograd engine calculate exact gradients for every variable automatically.
The Everyday Robot Problem
Suppose you are writing an algorithm to optimize your robot’s motor acceleration, arm feedforward voltage, and PID gains simultaneously.
Writing out 20 pages of manual calculus derivatives by hand is tedious and error-prone. What if your code could record every math operation you perform, build a graph in the background, and calculate all partial derivatives automatically with a single call to loss.backward()?
1. The 30-Line Value Class
To build an autograd engine, we wrap raw floating-point numbers in a Value class that stores:
data: The scalar value (e.g.3.0).grad: The derivative of the final output with respect to this value (starts at0.0)._prev: The child nodes that produced this value._backward: A tiny function that applies the local Chain Rule.
class Value:
def __init__(self, data, _children=()):
self.data = float(data)
self.grad = 0.0
self._prev = set(_children)
self._backward = lambda: None
def __add__(self, other):
other = other if isinstance(other, Value) else Value(other)
out = Value(self.data + other.data, (self, other))
def _backward():
self.grad += 1.0 * out.grad
other.grad += 1.0 * out.grad
out._backward = _backward
return out
def __mul__(self, other):
other = other if isinstance(other, Value) else Value(other)
out = Value(self.data * other.data, (self, other))
def _backward():
self.grad += other.data * out.grad
other.grad += self.data * out.grad
out._backward = _backward
return out
def relu(self):
out = Value(max(0.0, self.data), (self,))
def _backward():
self.grad += (1.0 if self.data > 0 else 0.0) * out.grad
out._backward = _backward
return out
def backward(self):
# 1. Build topological order of all nodes in the DAG
topo = []
visited = set()
def build_topo(v):
if v not in visited:
visited.add(v)
for child in v._prev:
build_topo(child)
topo.append(v)
build_topo(self)
# 2. Base gradient: d(Out) / d(Out) = 1.0
self.grad = 1.0
# 3. Traverse in reverse order and propagate gradients!
for node in reversed(topo):
node._backward()
2. Solving It in Code (Java Micro-Autograd Engine)
Here is a complete, self-contained Value class with automatic differentiation in pure Java:
import java.util.*;
public class Value {
public double data;
public double grad = 0.0;
private final List<Value> prev;
private Runnable backward = () -> {};
public Value(double data, Value... children) {
this.data = data;
this.prev = Arrays.asList(children);
}
public Value add(Value other) {
Value out = new Value(this.data + other.data, this, other);
out.backward = () -> {
this.grad += 1.0 * out.grad;
other.grad += 1.0 * out.grad;
};
return out;
}
public Value mul(Value other) {
Value out = new Value(this.data * other.data, this, other);
out.backward = () -> {
this.grad += other.data * out.grad;
other.grad += this.data * out.grad;
};
return out;
}
public void backward() {
List<Value> topo = new ArrayList<>();
Set<Value> visited = new HashSet<>();
buildTopo(this, topo, visited);
this.grad = 1.0;
for (int i = topo.size() - 1; i >= 0; i--) {
topo.get(i).backward.run();
}
}
private void buildTopo(Value v, List<Value> topo, Set<Value> visited) {
if (!visited.contains(v)) {
visited.add(v);
for (Value child : v.prev) buildTopo(child, topo, visited);
topo.add(v);
}
}
public static void main(String[] args) {
Value x = new Value(2.0);
Value w = new Value(3.0);
Value b = new Value(1.0);
// Forward: y = w * x + b
Value y = w.mul(x).add(b); // 7.0
// Loss: L = (y - 10)^2
Value diff = y.add(new Value(-10.0));
Value loss = diff.mul(diff); // 9.0
// Auto-differentiate!
loss.backward();
System.out.printf("Loss: %.2f | dL/dw: %.2f | dL/dx: %.2f | dL/db: %.2f%n",
loss.data, w.grad, x.grad, b.grad);
}
}
3. Math! Translation Sidebar
Why self.grad += ... Instead of self.grad = ...?
When a single variable is used in more than one place (for example: f = x * x), the multivariable Chain Rule states that gradients from all branches add together:
dL / dx = ∑ (dL / d_branchᵢ) · (d_branchᵢ / dx)
Using += ensures that when multiple operations reuse the same weight, all gradient paths accumulate correctly without overwriting each other!
4. Bridge to PyTorch & Deep Learning
torch.Tensor: PyTorch operates identically to ourValueclass, but instead of single scalars, it executes operations on multi-dimensional matrices and tensors in parallel on GPUs.- Training Loop: Every deep learning training step in PyTorch follows the exact same 3-step rhythm:
loss = model(inputs)(Forward Pass)loss.backward()(Backprop through computational graph)optimizer.step()(Gradient descent update)