use crate::{Node, NodeType, collections::GraphNode};
use radiate_core::Valid;
use std::collections::VecDeque;
pub trait GraphIterator<'a, T> {
fn iter_topological(&'a self) -> GraphTopologicalIterator<'a, T>;
fn get_nodes_of_type(
&'a self,
node_type: NodeType,
) -> impl Iterator<Item = &'a GraphNode<T>> + 'a
where
Self: AsRef<[GraphNode<T>]>,
T: 'a,
{
self.as_ref()
.iter()
.filter(move |node| node.node_type() == node_type)
}
}
impl<'a, G: AsRef<[GraphNode<T>]>, T> GraphIterator<'a, T> for G {
fn iter_topological(&'a self) -> GraphTopologicalIterator<'a, T> {
GraphTopologicalIterator::new(self.as_ref())
}
}
pub struct GraphTopologicalIterator<'a, T> {
graph: &'a [GraphNode<T>],
completed: Vec<bool>,
index_queue: VecDeque<usize>,
pending_index: usize,
remaining: usize,
}
impl<'a, T> GraphTopologicalIterator<'a, T> {
pub fn new(graph: &'a [GraphNode<T>]) -> Self {
let is_valid = !graph.iter().any(|node| !node.is_valid());
GraphTopologicalIterator {
graph,
completed: vec![false; graph.len()],
index_queue: VecDeque::with_capacity(graph.len()),
pending_index: if is_valid { 0 } else { graph.len() },
remaining: if is_valid { graph.len() } else { 0 },
}
}
}
impl<'a, T> Iterator for GraphTopologicalIterator<'a, T> {
type Item = &'a GraphNode<T>;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
let mut min_pending_index = self.graph.len();
for index in self.pending_index..self.graph.len() {
if self.completed[index] {
continue;
}
let node = &self.graph[index];
let mut degree = node.incoming().len();
for incoming_index in node.incoming() {
let incoming_node = &self.graph[*incoming_index];
if self.completed[incoming_node.index()] || incoming_node.is_recurrent() {
degree -= 1;
}
}
if degree == 0 {
self.completed[node.index()] = true;
self.index_queue.push_back(node.index());
} else {
min_pending_index = std::cmp::min(min_pending_index, node.index());
}
}
self.pending_index = min_pending_index;
self.index_queue.pop_front().map(|idx| {
self.remaining -= 1;
&self.graph[idx]
})
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.remaining, Some(self.remaining))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::collections::{Graph, GraphNode, NodeType};
use crate::ops::Op;
#[test]
fn test_graph_iterator() {
let graph = Graph::<Op<f64>>::new(vec![
GraphNode::from((0, NodeType::Input, Op::var(0))).with_outgoing([2]),
GraphNode::from((1, NodeType::Input, Op::var(1))).with_outgoing([2]),
GraphNode::from((2, NodeType::Vertex, Op::add()))
.with_incoming([0, 1])
.with_outgoing([3]),
GraphNode::from((3, NodeType::Output, Op::linear())).with_incoming([2]),
]);
let mut iter = graph.iter_topological();
assert_eq!(iter.next().unwrap().index(), 0);
assert_eq!(iter.next().unwrap().index(), 1);
assert_eq!(iter.next().unwrap().index(), 2);
assert_eq!(iter.next().unwrap().index(), 3);
assert!(iter.next().is_none());
}
#[test]
fn test_graph_iterator_recurrent() {
let nodes = vec![
GraphNode::from((0, NodeType::Input, Op::<f64>::var(0), vec![], vec![2])),
GraphNode::from((1, NodeType::Input, Op::<f64>::var(1), vec![], vec![2])),
GraphNode::from((2, NodeType::Vertex, Op::<f64>::add(), vec![0, 1], vec![3])),
GraphNode::from((3, NodeType::Vertex, Op::<f64>::mul(), vec![2], vec![2])),
GraphNode::from((4, NodeType::Output, Op::<f64>::linear(), vec![3], vec![])),
];
let graph = Graph::new(nodes);
let mut iter = graph.iter_topological();
assert_eq!(iter.next().unwrap().index(), 0);
assert_eq!(iter.next().unwrap().index(), 1);
assert_eq!(iter.next().unwrap().index(), 2);
assert_eq!(iter.next().unwrap().index(), 3);
assert_eq!(iter.next().unwrap().index(), 4);
assert!(iter.next().is_none());
}
#[test]
fn test_graph_iterator_disconnected() {
let nodes = vec![
GraphNode::from((0, NodeType::Input, Op::<f64>::var(0))).with_outgoing([2]),
GraphNode::from((1, NodeType::Input, Op::<f64>::var(1))),
GraphNode::from((2, NodeType::Vertex, Op::<f64>::add())).with_incoming([0]),
GraphNode::from((3, NodeType::Output, Op::<f64>::linear())).with_incoming([2]),
];
let results = Graph::new(nodes)
.iter_topological()
.map(|node| node.index())
.collect::<Vec<usize>>();
assert!(results.is_empty());
}
#[test]
fn test_graph_deep_cycles() {
let mut graph = Graph::<Op<f32>>::default();
graph.insert(NodeType::Input, Op::var(0));
graph.insert(NodeType::Vertex, Op::diff());
graph.insert(NodeType::Output, Op::sigmoid());
graph.insert(NodeType::Vertex, Op::div());
graph.insert(NodeType::Vertex, Op::pow());
graph.insert(NodeType::Edge, Op::weight());
graph.insert(NodeType::Edge, Op::identity());
graph.insert(NodeType::Vertex, Op::exp());
graph.insert(NodeType::Vertex, Op::cos());
graph.insert(NodeType::Edge, Op::weight());
graph.attach(0, 1);
graph.attach(1, 1);
graph.attach(4, 1);
graph.attach(7, 1);
graph.attach(1, 2);
graph.attach(3, 2);
graph.attach(9, 2);
graph.attach(0, 3);
graph.attach(5, 3);
graph.attach(0, 4);
graph.attach(8, 4);
graph.attach(1, 5);
graph.attach(3, 6);
graph.attach(4, 7);
graph.attach(6, 8);
graph.attach(7, 9);
graph.set_cycles(vec![]);
let results = graph
.iter_topological()
.map(|node| node.index())
.collect::<Vec<usize>>();
assert_eq!(results, vec![0, 1, 3, 4, 5, 6, 7, 8, 9, 2]);
}
#[test]
fn test_size_hint_tracks_actual_remaining() {
let graph = Graph::<Op<f64>>::new(vec![
GraphNode::from((0, NodeType::Input, Op::var(0))).with_outgoing([2]),
GraphNode::from((1, NodeType::Input, Op::var(1))).with_outgoing([2]),
GraphNode::from((2, NodeType::Vertex, Op::add()))
.with_incoming([0, 1])
.with_outgoing([3]),
GraphNode::from((3, NodeType::Output, Op::linear())).with_incoming([2]),
]);
let mut iter = graph.iter_topological();
let mut actual_remaining = 4;
assert_eq!(iter.size_hint(), (actual_remaining, Some(actual_remaining)));
iter.next();
actual_remaining -= 1;
assert_eq!(iter.size_hint(), (actual_remaining, Some(actual_remaining)));
while iter.next().is_some() {
actual_remaining -= 1;
assert_eq!(iter.size_hint(), (actual_remaining, Some(actual_remaining)));
}
assert_eq!(actual_remaining, 0);
assert_eq!(iter.size_hint(), (0, Some(0)));
}
#[test]
fn test_invalid_graph_size_hint_is_zero() {
let nodes = vec![
GraphNode::from((0, NodeType::Input, Op::<f64>::var(0))).with_outgoing([2]),
GraphNode::from((1, NodeType::Input, Op::<f64>::var(1))),
GraphNode::from((2, NodeType::Vertex, Op::<f64>::add())).with_incoming([0]),
GraphNode::from((3, NodeType::Output, Op::<f64>::linear())).with_incoming([2]),
];
let graph = Graph::new(nodes);
let iter = graph.iter_topological();
assert_eq!(iter.size_hint(), (0, Some(0)));
}
#[test]
fn test_eval_order_collect_allocates_exactly_once() {
let mut graph = Graph::<Op<f32>>::default();
graph.insert(NodeType::Input, Op::var(0));
graph.insert(NodeType::Vertex, Op::diff());
graph.insert(NodeType::Output, Op::sigmoid());
graph.insert(NodeType::Vertex, Op::div());
graph.insert(NodeType::Vertex, Op::pow());
graph.insert(NodeType::Edge, Op::weight());
graph.insert(NodeType::Edge, Op::identity());
graph.insert(NodeType::Vertex, Op::exp());
graph.insert(NodeType::Vertex, Op::cos());
graph.insert(NodeType::Edge, Op::weight());
graph.attach(0, 1);
graph.attach(1, 1);
graph.attach(4, 1);
graph.attach(7, 1);
graph.attach(1, 2);
graph.attach(3, 2);
graph.attach(9, 2);
graph.attach(0, 3);
graph.attach(5, 3);
graph.attach(0, 4);
graph.attach(8, 4);
graph.attach(1, 5);
graph.attach(3, 6);
graph.attach(4, 7);
graph.attach(6, 8);
graph.attach(7, 9);
graph.set_cycles(vec![]);
let eval_order = graph
.iter_topological()
.map(|n| n.index())
.collect::<Vec<usize>>();
assert_eq!(eval_order.len(), graph.len());
assert_eq!(eval_order.capacity(), eval_order.len());
}
}