use std::cell::RefCell;
const NONE: usize = usize::MAX;
#[derive(Debug, Clone, Copy)]
pub(crate) struct Node {
pub(crate) parents: [usize; 2],
pub(crate) partials: [f64; 2],
}
#[derive(Debug, Default)]
pub struct Tape {
pub(crate) nodes: RefCell<Vec<Node>>,
}
impl Tape {
pub fn new() -> Tape {
Tape::default()
}
pub fn var(&self, value: f64) -> super::var::Var<'_> {
let idx = self.push0();
super::var::Var { tape: self, idx, val: value }
}
pub fn len(&self) -> usize {
self.nodes.borrow().len()
}
pub fn is_empty(&self) -> bool {
self.nodes.borrow().is_empty()
}
pub fn clear(&self) {
self.nodes.borrow_mut().clear();
}
pub(crate) fn push0(&self) -> usize {
let mut nodes = self.nodes.borrow_mut();
nodes.push(Node { parents: [NONE, NONE], partials: [0.0, 0.0] });
nodes.len() - 1
}
pub(crate) fn push1(&self, parent: usize, partial: f64) -> usize {
let mut nodes = self.nodes.borrow_mut();
nodes.push(Node { parents: [parent, NONE], partials: [partial, 0.0] });
nodes.len() - 1
}
pub(crate) fn push2(&self, p0: usize, w0: f64, p1: usize, w1: f64) -> usize {
let mut nodes = self.nodes.borrow_mut();
nodes.push(Node { parents: [p0, p1], partials: [w0, w1] });
nodes.len() - 1
}
pub(crate) fn backward(&self, output: usize) -> Vec<f64> {
let nodes = self.nodes.borrow();
let mut adjoint = vec![0.0; nodes.len()];
adjoint[output] = 1.0;
for i in (0..=output).rev() {
let a = adjoint[i];
if a == 0.0 {
continue;
}
let node = nodes[i];
for slot in 0..2 {
let p = node.parents[slot];
if p != NONE {
adjoint[p] += node.partials[slot] * a;
}
}
}
adjoint
}
}
#[derive(Debug, Clone)]
pub struct Gradients {
pub(crate) adjoints: Vec<f64>,
}
impl Gradients {
pub fn wrt(&self, v: super::var::Var<'_>) -> f64 {
self.adjoints[v.idx]
}
}