use std::collections::VecDeque;
use crate::error::{GraphError, GraphResult};
use crate::graph::ComputeGraph;
use crate::node::NodeId;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Wavefront {
pub level: usize,
pub nodes: Vec<NodeId>,
pub max_cost: u64,
pub total_cost: u64,
}
impl Wavefront {
#[must_use]
pub fn width(&self) -> usize {
self.nodes.len()
}
}
#[derive(Debug, Clone)]
pub struct Schedule {
waves: Vec<Wavefront>,
levels: Vec<usize>,
critical_path_cost: u64,
}
impl Schedule {
pub fn levelize(graph: &ComputeGraph) -> GraphResult<Self> {
if graph.is_empty() {
return Err(GraphError::EmptyGraph);
}
let n = graph.node_count();
let mut levels = vec![0usize; n];
let mut in_degree: Vec<u32> = (0..n)
.map(|i| {
graph
.predecessors(NodeId(i as u32))
.map(|p| p.len() as u32)
.unwrap_or(0)
})
.collect();
let mut queue: VecDeque<NodeId> = (0..n)
.filter(|&i| in_degree[i] == 0)
.map(|i| NodeId(i as u32))
.collect();
let mut processed = 0usize;
while let Some(id) = queue.pop_front() {
processed += 1;
let lv = levels[id.0 as usize];
for &succ in graph.successors(id)? {
let nl = lv + 1;
if nl > levels[succ.0 as usize] {
levels[succ.0 as usize] = nl;
}
let d = &mut in_degree[succ.0 as usize];
*d -= 1;
if *d == 0 {
queue.push_back(succ);
}
}
}
debug_assert_eq!(processed, n, "levelization did not visit every node");
let max_level = *levels.iter().max().unwrap_or(&0);
let mut buckets: Vec<Vec<NodeId>> = vec![Vec::new(); max_level + 1];
for i in 0..n {
buckets[levels[i]].push(NodeId(i as u32));
}
let order = graph.topological_order()?;
let mut dist = vec![0u64; n];
for &id in &order {
let cost = graph.node(id)?.cost_hint;
let mut best_pred = 0u64;
for &pred in graph.predecessors(id)? {
best_pred = best_pred.max(dist[pred.0 as usize]);
}
dist[id.0 as usize] = best_pred + cost;
}
let critical_path_cost = dist.iter().copied().max().unwrap_or(0);
let waves: Vec<Wavefront> = buckets
.into_iter()
.enumerate()
.map(|(level, mut nodes)| {
nodes.sort();
let max_cost = nodes
.iter()
.map(|&id| graph.nodes()[id.0 as usize].cost_hint)
.max()
.unwrap_or(0);
let total_cost: u64 = nodes
.iter()
.map(|&id| graph.nodes()[id.0 as usize].cost_hint)
.sum();
Wavefront {
level,
nodes,
max_cost,
total_cost,
}
})
.collect();
Ok(Self {
waves,
levels,
critical_path_cost,
})
}
#[must_use]
pub fn wavefronts(&self) -> &[Wavefront] {
&self.waves
}
#[must_use]
pub fn depth(&self) -> usize {
self.waves.len()
}
pub fn level_of(&self, id: NodeId) -> GraphResult<usize> {
self.levels
.get(id.0 as usize)
.copied()
.ok_or(GraphError::NodeNotFound(id))
}
#[must_use]
pub fn max_width(&self) -> usize {
self.waves.iter().map(Wavefront::width).max().unwrap_or(0)
}
#[must_use]
pub fn critical_path_cost(&self) -> u64 {
self.critical_path_cost
}
#[must_use]
pub fn bounded_makespan(&self, max_streams: usize) -> u64 {
let lanes = max_streams.max(1);
self.waves.iter().map(|w| wave_makespan(w, lanes)).sum()
}
#[must_use]
pub fn unbounded_makespan(&self) -> u64 {
self.waves.iter().map(|w| w.max_cost).sum()
}
}
fn wave_makespan(wave: &Wavefront, lanes: usize) -> u64 {
if wave.nodes.is_empty() {
return 0;
}
if lanes >= wave.nodes.len() {
return wave.max_cost;
}
let balanced = wave.total_cost.div_ceil(lanes as u64);
wave.max_cost.max(balanced)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::builder::GraphBuilder;
fn cost_node(b: &mut GraphBuilder, name: &str, cost: u64) -> NodeId {
b.add_raw(
crate::node::GraphNode::new(NodeId(0), crate::node::NodeKind::Barrier)
.with_name(name)
.with_cost(cost),
)
}
#[test]
fn levelize_empty_errors() {
let g = ComputeGraph::new();
assert!(matches!(
Schedule::levelize(&g),
Err(GraphError::EmptyGraph)
));
}
#[test]
fn linear_chain_one_node_per_wave() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = b.add_barrier("a");
let c = b.add_barrier("b");
let d = b.add_barrier("c");
b.chain(&[a, c, d]);
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
assert_eq!(sch.depth(), 3);
for w in sch.wavefronts() {
assert_eq!(w.width(), 1);
}
assert_eq!(sch.max_width(), 1);
}
#[test]
fn fork_join_groups_independent_nodes() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let src = b.add_barrier("src");
let a = b.add_barrier("a");
let bb = b.add_barrier("b");
let c = b.add_barrier("c");
let sink = b.add_barrier("sink");
b.fan_out(src, &[a, bb, c]);
b.fan_in(&[a, bb, c], sink);
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
assert_eq!(sch.depth(), 3);
let mid = &sch.wavefronts()[1];
assert_eq!(mid.width(), 3);
let mut got = mid.nodes.clone();
got.sort();
let mut want = vec![a, bb, c];
want.sort();
assert_eq!(got, want);
assert_eq!(sch.max_width(), 3);
}
#[test]
fn wavefront_nodes_are_mutually_independent() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = b.add_barrier("a");
let bb = b.add_barrier("b");
let c = b.add_barrier("c");
let d = b.add_barrier("d");
let e = b.add_barrier("e");
b.dep(a, c);
b.dep(a, d);
b.dep(bb, d);
b.dep(bb, e);
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
for wave in sch.wavefronts() {
for &u in &wave.nodes {
for &v in &wave.nodes {
if u != v {
assert!(
!g.is_reachable(u, v),
"wave-mates {u} and {v} must be independent"
);
}
}
}
}
}
#[test]
fn levels_match_longest_path() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = b.add_barrier("a");
let bb = b.add_barrier("b");
let c = b.add_barrier("c");
let d = b.add_barrier("d");
b.dep(a, bb).dep(a, c).dep(bb, d).dep(c, d);
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
assert_eq!(sch.level_of(a).expect("a"), 0);
assert_eq!(sch.level_of(bb).expect("b"), 1);
assert_eq!(sch.level_of(c).expect("c"), 1);
assert_eq!(sch.level_of(d).expect("d"), 2);
}
#[test]
fn level_of_out_of_range() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
b.add_barrier("a");
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
assert!(matches!(
sch.level_of(NodeId(50)),
Err(GraphError::NodeNotFound(_))
));
}
#[test]
fn critical_path_cost_weighted() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = cost_node(&mut b, "a", 1);
let bb = cost_node(&mut b, "b", 10);
let c = cost_node(&mut b, "c", 1);
b.chain(&[a, bb, c]);
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
assert_eq!(sch.critical_path_cost(), 12);
}
#[test]
fn critical_path_takes_longest_branch() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = cost_node(&mut b, "a", 1);
let bb = cost_node(&mut b, "b", 5);
let c = cost_node(&mut b, "c", 20);
let d = cost_node(&mut b, "d", 1);
b.dep(a, bb).dep(a, c).dep(bb, d).dep(c, d);
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
assert_eq!(sch.critical_path_cost(), 22);
}
#[test]
fn bounded_makespan_serializes_wide_wave() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let src = cost_node(&mut b, "src", 0);
let leaves: Vec<NodeId> = (0..4)
.map(|i| cost_node(&mut b, &format!("l{i}"), 10))
.collect();
b.fan_out(src, &leaves);
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
assert_eq!(sch.unbounded_makespan(), 10);
assert_eq!(sch.bounded_makespan(4), 10);
assert_eq!(sch.bounded_makespan(2), 20);
assert_eq!(sch.bounded_makespan(1), 40);
assert_eq!(sch.bounded_makespan(0), 40);
}
#[test]
fn bounded_makespan_respects_max_cost_lower_bound() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let src = cost_node(&mut b, "src", 0);
let big = cost_node(&mut b, "big", 30);
let s1 = cost_node(&mut b, "s1", 1);
let s2 = cost_node(&mut b, "s2", 1);
let s3 = cost_node(&mut b, "s3", 1);
b.fan_out(src, &[big, s1, s2, s3]);
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
assert_eq!(sch.bounded_makespan(2), 30);
}
#[test]
fn wave_cost_aggregates() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let src = cost_node(&mut b, "src", 0);
let a = cost_node(&mut b, "a", 3);
let c = cost_node(&mut b, "c", 7);
b.fan_out(src, &[a, c]);
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
let wave1 = &sch.wavefronts()[1];
assert_eq!(wave1.max_cost, 7);
assert_eq!(wave1.total_cost, 10);
}
#[test]
fn isolated_nodes_all_in_wave_zero() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
b.add_barrier("a");
b.add_barrier("b");
b.add_barrier("c");
let g = b.build().expect("builds");
let sch = Schedule::levelize(&g).expect("levelize");
assert_eq!(sch.depth(), 1);
assert_eq!(sch.wavefronts()[0].width(), 3);
}
}