use std::collections::HashSet;
use crate::tensor::backend::Backend;
use crate::tensor::graph::{NodeKind, TensorGraphNode};
use crate::tensor::planner::get_id;
pub(crate) struct TopologicalSortIter<'a, T, B: Backend> {
stack: Vec<(&'a NodeKind<T, B>, bool)>,
visited: HashSet<usize>,
}
impl<'a, T, B: Backend> TopologicalSortIter<'a, T, B> {
pub(crate) fn new(base_node: &'a TensorGraphNode<T, B>) -> Self {
let mut stack = Vec::new();
stack.extend(base_node.inputs.iter().map(|i| (i, false)));
Self {
stack,
visited: HashSet::new(),
}
}
}
impl<'a, T, B: Backend> Iterator for TopologicalSortIter<'a, T, B> {
type Item = &'a NodeKind<T, B>;
fn next(&mut self) -> Option<Self::Item> {
loop {
let (node, exiting) = self.stack.pop()?;
if exiting {
return Some(node);
}
let id = get_id(node);
if !self.visited.insert(id) {
continue;
}
self.stack.push((node, true));
match node {
NodeKind::Edge(_) | NodeKind::Slot(_) => {}
NodeKind::Node(n) => self.stack.extend(n.inputs.iter().map(|i| (i, false))),
NodeKind::Cache(cache) => {
if !cache.is_cache_filled() {
self.stack
.extend(cache.get_node().inputs.iter().rev().map(|i| (i, false)))
}
}
NodeKind::Baked(baked) => {
self.stack.extend(baked.inputs.iter().map(|i| (i, false)));
}
}
}
}
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(
level = "trace",
skip(base_node),
fields(node_id = base_node.id, inputs_count = base_node.inputs.len())
)
)]
pub(crate) fn topological_sort<T, B: Backend>(
base_node: &TensorGraphNode<T, B>,
) -> TopologicalSortIter<'_, T, B> {
TopologicalSortIter::new(base_node)
}