use std::ops::{Add, Div, Mul, Neg, Sub};
use topos::{Detach, Differentiable, Element, Elementary, Tape};
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Dual {
pub primal: f64,
pub tangent: f64,
}
impl Dual {
pub fn value(primal: f64) -> Self {
Self {
primal,
tangent: 0.0,
}
}
pub fn var(primal: f64) -> Self {
Self {
primal,
tangent: 1.0,
}
}
}
impl Add for Dual {
type Output = Self;
fn add(self, rhs: Self) -> Self {
Self {
primal: self.primal + rhs.primal,
tangent: self.tangent + rhs.tangent,
}
}
}
impl Sub for Dual {
type Output = Self;
fn sub(self, rhs: Self) -> Self {
Self {
primal: self.primal - rhs.primal,
tangent: self.tangent - rhs.tangent,
}
}
}
impl Mul for Dual {
type Output = Self;
fn mul(self, rhs: Self) -> Self {
Self {
primal: self.primal * rhs.primal,
tangent: self.primal * rhs.tangent + self.tangent * rhs.primal,
}
}
}
impl Div for Dual {
type Output = Self;
fn div(self, rhs: Self) -> Self {
let primal = self.primal / rhs.primal;
Self {
tangent: (self.tangent - primal * rhs.tangent) / rhs.primal,
primal,
}
}
}
impl Neg for Dual {
type Output = Self;
fn neg(self) -> Self {
Self {
primal: -self.primal,
tangent: -self.tangent,
}
}
}
impl Differentiable for Dual {
type Accumulator = Self;
fn promote(&self) -> Self {
*self
}
fn demote(accumulated: Self) -> Self {
accumulated
}
fn zero() -> Self {
Self::value(0.0)
}
fn one() -> Self {
Self::value(1.0)
}
fn from_count(count: usize) -> Self {
Self::value(count as f64)
}
fn is_count(&self, count: usize) -> bool {
self.primal == count as f64 && self.tangent == 0.0
}
}
impl Elementary for Dual {
fn exp(&self) -> Self {
let primal = Elementary::exp(&self.primal);
Self {
tangent: primal * self.tangent,
primal,
}
}
fn ln(&self) -> Self {
Self {
primal: Elementary::ln(&self.primal),
tangent: self.tangent / self.primal,
}
}
fn sqrt(&self) -> Self {
let primal = Elementary::sqrt(&self.primal);
Self {
tangent: self.tangent / (primal + primal),
primal,
}
}
fn tanh(&self) -> Self {
let primal = Elementary::tanh(&self.primal);
Self {
tangent: self.tangent * (1.0 - primal * primal),
primal,
}
}
fn sin(&self) -> Self {
Self {
primal: Elementary::sin(&self.primal),
tangent: self.tangent * Elementary::cos(&self.primal),
}
}
fn cos(&self) -> Self {
Self {
primal: Elementary::cos(&self.primal),
tangent: -(self.tangent * Elementary::sin(&self.primal)),
}
}
fn log1p(&self) -> Self {
Self {
primal: Elementary::log1p(&self.primal),
tangent: self.tangent / (1.0 + self.primal),
}
}
fn expm1(&self) -> Self {
let primal = Elementary::expm1(&self.primal);
Self {
tangent: self.tangent * (primal + 1.0),
primal,
}
}
fn erf(&self) -> Self {
Self {
primal: Elementary::erf(&self.primal),
tangent: self.tangent * Elementary::erf_derivative(&self.primal),
}
}
fn erf_derivative(&self) -> Self {
let primal = Elementary::erf_derivative(&self.primal);
Self {
tangent: self.tangent * (-2.0 * self.primal) * primal,
primal,
}
}
fn powf(&self, exponent: Self) -> Self {
let primal = Elementary::powf(&self.primal, exponent.primal);
Self {
tangent: primal
* (exponent.tangent * Elementary::ln(&self.primal)
+ exponent.primal * self.tangent / self.primal),
primal,
}
}
fn maximum(&self, other: &Self) -> Self {
Self {
primal: Elementary::maximum(&self.primal, &other.primal),
tangent: if self.primal >= other.primal {
self.tangent
} else {
other.tangent
},
}
}
fn step(&self, threshold: &Self) -> Self {
Self::value(Elementary::step(&self.primal, &threshold.primal))
}
}
impl Element for Dual {}
fn main() {
for (w0, x0, y0) in [(0.5, 0.25, 1.5), (-0.75, 1.25, 0.5), (2.0, -0.5, -1.25)] {
let tape: Tape<f64> = Tape::new();
let w = tape.parameter(w0);
let x = tape.input(x0);
let y = tape.input(y0);
let error = w * x - y;
let loss = error * error;
let adjoints = tape.differentiate(loss, [w]);
let (w, loss) = (w.symbol(), loss.symbol());
let network = tape.into_network();
let run = network.forward(&network.parameters(), []);
let engine = run.backward(loss).of(w).scalar();
let recorded = run.of(adjoints.pairs()[0].1).scalar();
let (dual_network, [dual_loss]) = Tape::record(|tape| {
let w = tape.parameter(Dual::var(w0));
let x = tape.input(Dual::value(x0));
let y = tape.input(Dual::value(y0));
let error = w * x - y;
[error * error].detach()
});
let dual_run = dual_network.forward(&dual_network.parameters(), []);
let slope = dual_run.of(dual_loss).scalar();
assert_eq!(
slope.tangent.to_bits(),
engine.to_bits(),
"the dual tangent must be the engine gradient, bit for bit"
);
assert_eq!(
slope.tangent.to_bits(),
recorded.to_bits(),
"the dual tangent must be the recorded gradient, bit for bit"
);
println!(
"d/dw (w*{x0} - {y0})^2 at w = {w0}: {} on all three routes",
slope.tangent
);
}
println!("forward mode by payload, reverse mode by scan: one answer");
}