use std::ops::{Add, Div, Mul, Neg, Sub};
use topos::{Detach, Network, Node, Numerics, Recordable, Shape, Symbol, Tape, Tensor, Trace};
#[derive(Clone, Debug)]
struct Dual<R> {
primal: R,
tangent: R,
}
impl<R: Recordable> Dual<R> {
fn constant(primal: R) -> Self {
let tangent = primal.zero_like();
Self { primal, tangent }
}
}
impl<R: Recordable> Add for Dual<R> {
type Output = Self;
fn add(self, rhs: Self) -> Self {
Self {
primal: self.primal + rhs.primal,
tangent: self.tangent + rhs.tangent,
}
}
}
impl<R: Recordable> Sub for Dual<R> {
type Output = Self;
fn sub(self, rhs: Self) -> Self {
Self {
primal: self.primal - rhs.primal,
tangent: self.tangent - rhs.tangent,
}
}
}
impl<R: Recordable> Mul for Dual<R> {
type Output = Self;
fn mul(self, rhs: Self) -> Self {
Self {
primal: self.primal.clone() * rhs.primal.clone(),
tangent: self.primal * rhs.tangent + self.tangent * rhs.primal,
}
}
}
impl<R: Recordable> Div for Dual<R> {
type Output = Self;
fn div(self, rhs: Self) -> Self {
let primal = self.primal / rhs.primal.clone();
let tangent = (self.tangent - primal.clone() * rhs.tangent) / rhs.primal;
Self { primal, tangent }
}
}
impl<R: Recordable> Neg for Dual<R> {
type Output = Self;
fn neg(self) -> Self {
Self {
primal: -self.primal,
tangent: -self.tangent,
}
}
}
impl<R: Recordable> Recordable for Dual<R> {
fn shape(&self) -> Shape {
self.primal.shape()
}
fn zero_like(&self) -> Self {
Self::constant(self.primal.zero_like())
}
fn one_like(&self) -> Self {
Self::constant(self.primal.one_like())
}
fn exp(&self) -> Self {
let primal = self.primal.exp();
let tangent = primal.clone() * self.tangent.clone();
Self { primal, tangent }
}
fn ln(&self) -> Self {
Self {
primal: self.primal.ln(),
tangent: self.tangent.clone() / self.primal.clone(),
}
}
fn sqrt(&self) -> Self {
let primal = self.primal.sqrt();
let tangent = self.tangent.clone() / (primal.clone() + primal.clone());
Self { primal, tangent }
}
fn tanh(&self) -> Self {
let primal = self.primal.tanh();
let tangent = self.tangent.clone() * (primal.one_like() - primal.clone() * primal.clone());
Self { primal, tangent }
}
fn sin(&self) -> Self {
Self {
primal: self.primal.sin(),
tangent: self.tangent.clone() * self.primal.cos(),
}
}
fn cos(&self) -> Self {
Self {
primal: self.primal.cos(),
tangent: -(self.tangent.clone() * self.primal.sin()),
}
}
fn log1p(&self) -> Self {
Self {
primal: self.primal.log1p(),
tangent: self.tangent.clone() / (self.primal.one_like() + self.primal.clone()),
}
}
fn expm1(&self) -> Self {
let primal = self.primal.expm1();
let tangent = self.tangent.clone() * (primal.clone() + primal.one_like());
Self { primal, tangent }
}
fn erf(&self) -> Self {
Self {
primal: self.primal.erf(),
tangent: self.tangent.clone() * self.primal.erf_derivative(),
}
}
fn erf_derivative(&self) -> Self {
let primal = self.primal.erf_derivative();
let tangent =
-((self.primal.clone() + self.primal.clone()) * primal.clone() * self.tangent.clone());
Self { primal, tangent }
}
fn powf(&self, exponent: Self) -> Self {
let primal = self.primal.powf(exponent.primal.clone());
let tangent = primal.clone()
* (exponent.tangent * self.primal.ln()
+ exponent.primal * (self.tangent.clone() / self.primal.clone()));
Self { primal, tangent }
}
fn maximum(&self, other: &Self) -> Self {
let winners = self.primal.step(&other.primal);
Self {
primal: self.primal.maximum(&other.primal),
tangent: winners.clone() * self.tangent.clone()
+ (winners.one_like() - winners) * other.tangent.clone(),
}
}
fn step(&self, threshold: &Self) -> Self {
Self::constant(self.primal.step(&threshold.primal))
}
fn matmul(&self, rhs: &Self) -> Self {
Self {
primal: self.primal.matmul(&rhs.primal),
tangent: self.tangent.matmul(&rhs.primal) + self.primal.matmul(&rhs.tangent),
}
}
fn sum(&self) -> Self {
Self {
primal: self.primal.sum(),
tangent: self.tangent.sum(),
}
}
fn sum_along(&self, axis: usize) -> Self {
Self {
primal: self.primal.sum_along(axis),
tangent: self.tangent.sum_along(axis),
}
}
fn logsumexp(&self, axis: usize) -> Self {
let primal = self.primal.logsumexp(axis);
let extent = self.primal.shape().axes()[axis];
let weights = (self.primal.clone() - primal.broadcast_along(axis, extent)).exp();
Self {
tangent: (weights * self.tangent.clone()).sum_along(axis),
primal,
}
}
fn log_softmax(&self, axis: usize) -> Self {
let primal = self.primal.log_softmax(axis);
let extent = self.primal.shape().axes()[axis];
let mass = self.tangent.sum_along(axis).broadcast_along(axis, extent);
Self {
tangent: self.tangent.clone() - primal.exp() * mass,
primal,
}
}
fn broadcast(&self, shape: Shape) -> Self {
Self {
primal: self.primal.broadcast(shape.clone()),
tangent: self.tangent.broadcast(shape),
}
}
fn broadcast_along(&self, axis: usize, extent: usize) -> Self {
Self {
primal: self.primal.broadcast_along(axis, extent),
tangent: self.tangent.broadcast_along(axis, extent),
}
}
fn reshape(&self, shape: Shape) -> Self {
Self {
primal: self.primal.reshape(shape.clone()),
tangent: self.tangent.reshape(shape),
}
}
fn permute(&self, order: &[usize]) -> Self {
Self {
primal: self.primal.permute(order),
tangent: self.tangent.permute(order),
}
}
fn narrow(&self, axis: usize, start: usize, len: usize) -> Self {
Self {
primal: self.primal.narrow(axis, start, len),
tangent: self.tangent.narrow(axis, start, len),
}
}
fn pad(&self, axis: usize, start: usize, full_extent: usize) -> Self {
Self {
primal: self.primal.pad(axis, start, full_extent),
tangent: self.tangent.pad(axis, start, full_extent),
}
}
fn unfold(&self, axis: usize, size: usize, step: usize, dilation: usize) -> Self {
Self {
primal: self.primal.unfold(axis, size, step, dilation),
tangent: self.tangent.unfold(axis, size, step, dilation),
}
}
fn fold(&self, axis: usize, size: usize, step: usize, dilation: usize, extent: usize) -> Self {
Self {
primal: self.primal.fold(axis, size, step, dilation, extent),
tangent: self.tangent.fold(axis, size, step, dilation, extent),
}
}
fn gather(&self, selection: &Self) -> Self {
Self {
primal: self.primal.gather(&selection.primal),
tangent: self.tangent.gather(&selection.primal),
}
}
fn scatter(&self, selection: &Self) -> Self {
Self {
primal: self.primal.scatter(&selection.primal),
tangent: self.tangent.scatter(&selection.primal),
}
}
}
fn jvp_walk<R: Recordable>(
nodes: &[Node],
mut source: impl FnMut(&Node) -> Dual<R>,
) -> Vec<Dual<R>> {
let mut duals: Vec<Dual<R>> = Vec::new();
for node in nodes {
let dual = if node.is_source() {
source(node)
} else {
let operands: Vec<&Dual<R>> = node
.operands()
.iter()
.map(|symbol| &duals[symbol.index()])
.collect();
node.opcode().express(&operands)
};
duals.push(dual);
}
duals
}
fn jvp_eager(network: &Network<f64>, wrt: Symbol, seed: &Tensor<f64>, read: Symbol) -> Tensor<f64> {
let nodes: Vec<Node> = network.nodes().collect();
let duals = Numerics::exactly(|| {
jvp_walk(&nodes, |node| {
let primal = network
.payload(node.symbol())
.expect("sources hold payloads")
.clone();
let tangent = if node.symbol() == wrt {
seed.clone()
} else {
primal.zero_like()
};
Dual { primal, tangent }
})
});
duals[read.index()].tangent.clone()
}
fn assert_bitwise(expected: &Tensor<f64>, computed: &Tensor<f64>, subject: &str) {
let expected = expected.to_vec();
let computed = computed.to_vec();
assert_eq!(expected.len(), computed.len(), "{subject}: length differs");
for (expected, computed) in expected.iter().zip(&computed) {
assert_eq!(
expected.to_bits(),
computed.to_bits(),
"{subject}: {computed} differs from {expected}"
);
}
}
fn main() {
let (network, [weights, _inputs, loss]) = Tape::record(|tape| {
let weights = tape.parameter(Tensor::new(
[3, 2],
vec![0.5, -0.25, 0.125, 0.75, -0.5, 0.25],
));
let inputs = tape.input(Tensor::new([2, 3], vec![0.5, -1.0, 0.25, -0.75, 0.5, 1.25]));
let scores = inputs.matmul(weights).tanh().log_softmax(1);
let loss = (scores * scores).sum();
[weights, inputs, loss].detach()
});
let seed = Tensor::new([3, 2], vec![0.25, -0.5, 1.0, 0.125, -0.25, 0.5]);
let eager = jvp_eager(&network, weights, &seed, loss);
let nodes: Vec<Node> = network.nodes().collect();
let tape = network.into_tape();
let duals = jvp_walk(&nodes, |node| {
let primal = Trace::of(tape.resolve(node.symbol()));
let payload = tape.payload(node.symbol()).expect("sources hold payloads");
let tangent_payload = if node.symbol() == weights {
seed.clone()
} else {
payload.zero_like()
};
let tangent = Trace::of(tape.leaf(tangent_payload));
Dual { primal, tangent }
});
let tangent_symbol = duals[loss.index()].tangent.value().symbol();
drop(duals);
let jvp_network = tape.into_network();
let recorded_run = jvp_network.forward(&jvp_network.parameters(), []);
assert_bitwise(
&eager,
recorded_run.of(tangent_symbol),
"recorded jvp against eager",
);
println!("eager and recorded jvp agree bitwise: {}", eager.scalar());
let (cubic, [x, value]) = Tape::record(|tape| {
let x = tape.parameter(Tensor::new([3], vec![0.5, -0.25, 1.5]));
let value = (x * x * x).sum();
[x, value].detach()
});
let direction = Tensor::new([3], vec![0.25, 1.0, -0.5]);
let forward_slope = jvp_eager(&cubic, x, &direction, value);
let run = cubic.forward(&cubic.parameters(), []);
let gradients = run.backward(value);
let reverse_slope = Numerics::exactly(|| (gradients.of(x).clone() * direction.clone()).sum());
assert_bitwise(&forward_slope, &reverse_slope, "slope against the gradient");
println!(
"forward slope equals <gradient, seed> bitwise: {}",
forward_slope.scalar()
);
let tape: Tape<f64> = Tape::new();
let x = tape.parameter(Tensor::new([3], vec![0.5, -0.25, 1.5]));
let value = (x * x * x).sum();
let adjoints = tape.differentiate(value, [x]);
let &(_, gradient_symbol) = &adjoints.pairs()[0];
let direction_leaf = tape.leaf(direction.clone());
let hessian_adjoints = tape.vjp(gradient_symbol, direction_leaf, [x]);
let &(_, reverse_hvp_symbol) = &hessian_adjoints.pairs()[0];
let x = x.symbol();
let network = tape.into_network();
let run = network.forward(&network.parameters(), []);
let reverse_hvp = run.of(reverse_hvp_symbol);
let forward_hvp = jvp_eager(&network, x, &direction, gradient_symbol);
assert_bitwise(
reverse_hvp,
&forward_hvp,
"forward-over-reverse against reverse-over-reverse",
);
println!(
"hessian-vector product agrees across modes: {:?}",
forward_hvp.to_vec()
);
}