use crate::tensor::Tensor;
use ndarray::ArrayD;
use std::collections::HashSet;
use std::fmt;
pub struct BackwardContext {
pub inputs: Vec<Tensor>,
pub backward_fn: Box<dyn Fn(&ArrayD<f32>)>,
}
impl fmt::Debug for BackwardContext {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BackwardContext")
.field("num_inputs", &self.inputs.len())
.finish()
}
}
pub fn backward(tensor: &Tensor) {
let mut sorted_graph = Vec::new();
let mut visited = HashSet::new();
fn build_graph(tensor: &Tensor, visited: &mut HashSet<Tensor>, sorted: &mut Vec<Tensor>) {
if visited.contains(tensor) {
return;
}
visited.insert(tensor.clone());
if let Some(ctx) = &tensor.ctx {
for input_tensor in &ctx.inputs {
build_graph(input_tensor, visited, sorted);
}
}
sorted.push(tensor.clone());
}
build_graph(tensor, &mut visited, &mut sorted_graph);
if let Some(grad) = &tensor.grad {
grad.borrow_mut().fill(1.0);
} else {
panic!("backward() called on a tensor that does not require gradients");
}
for t in sorted_graph.iter().rev() {
if let Some(ctx) = &t.ctx {
let upstream_grad = t.grad.as_ref().unwrap().borrow();
(ctx.backward_fn)(&upstream_grad);
}
}
}