use crate::{variable::Variable, Graph};
pub fn graphviz(graph: &Graph, vars: &[Variable]) -> String {
let mut dot = String::from(
"digraph Graph {\n\tbgcolor=\"transparent\";\n\trankdir=\"LR\";\n\tnode [shape=box3d];\n",
);
let vertices = graph.vertices.borrow();
let var_indices: std::collections::HashSet<_> = vars.iter().map(|var| var.index).collect();
for (index, _vertex) in vertices.iter().enumerate() {
if var_indices.contains(&index) {
let var_value = vars.iter().find(|var| var.index == index).unwrap().value();
dot.push_str(&format!(
"\t{} [label=\"Input: x_{}, Value: {:.2}\", color=\"red\"];\n",
index, index, var_value
));
} else {
dot.push_str(&format!("\t{} [label=\"Op: #{}\"];\n", index, index));
}
}
for (index, vertex) in vertices.iter().enumerate() {
for (i, parent) in vertex.parents.iter().enumerate() {
if parent != &index {
let label = vertex.partials[i];
dot.push_str(&format!(
"\t{} -> {} [label=\"\u{2202}_{}: {:.2?}\"];\n",
parent, index, i, label
));
}
}
}
dot.push_str("}\n");
dot
}
#[cfg(test)]
mod test_graphviz {
use super::*;
use crate::Powf;
use crate::{Accumulate, Gradient};
#[test]
fn test_graphviz_1() {
let graph = Graph::new();
let x = graph.var(2.0);
let y = graph.var(3.0);
let z = x * y;
let _u = (z.exp()).sin();
print!("{}", graphviz(&graph, &[x, y]));
}
#[test]
fn test_graphviz_2() {
let graph = Graph::new();
let a = graph.var(1.0);
let b = graph.var(2.0);
let c = graph.var(3.0);
let d = graph.var(4.0);
let f1 = (a.exp() + b.cbrt()).sin();
let e = graph.var(5.0);
let f2 = e * (c.sqrt() + d.powf(2.0)).ln();
let f3 = f1 + f2;
print!("{}", graphviz(&graph, &[a, b, c, d, e]));
println!("Gradient: {:.4?}", f3.accumulate().wrt(&[a, b, c, d, e]));
}
#[test]
fn test_graphviz_3() {
let graph = Graph::new();
let x = graph.var(1.0);
let y = graph.var(2.0);
let z = x * y + y.sin();
let _g = z.accumulate();
print!("{}", graphviz(&graph, &[x, y]));
}
}