use std::collections::HashMap;
use crate::analysis::topo_analyse;
use crate::error::{GraphError, GraphResult};
use crate::graph::ComputeGraph;
use crate::node::{NodeId, StreamId};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StreamAssignment {
pub node: NodeId,
pub stream: StreamId,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SyncPoint {
pub from: NodeId,
pub stream_from: StreamId,
pub to: NodeId,
pub stream_to: StreamId,
}
#[derive(Debug, Clone)]
pub struct StreamPlan {
pub assignments: Vec<StreamAssignment>,
pub sync_points: Vec<SyncPoint>,
pub num_streams: usize,
pub max_streams: usize,
}
impl StreamPlan {
pub fn stream_of(&self, node: NodeId) -> StreamId {
self.assignments
.iter()
.find(|a| a.node == node)
.map(|a| a.stream)
.unwrap_or(StreamId::DEFAULT)
}
pub fn nodes_on(&self, stream: StreamId) -> Vec<NodeId> {
self.assignments
.iter()
.filter(|a| a.stream == stream)
.map(|a| a.node)
.collect()
}
pub fn has_concurrency(&self) -> bool {
let first = self.assignments.first().map(|a| a.stream);
self.assignments.iter().any(|a| Some(a.stream) != first)
}
pub fn sync_count(&self) -> usize {
self.sync_points.len()
}
}
pub fn analyse(graph: &ComputeGraph, max_streams: usize) -> GraphResult<StreamPlan> {
if graph.is_empty() {
return Err(GraphError::EmptyGraph);
}
let max_streams = max_streams.max(1);
let topo = topo_analyse(graph)?;
let priority = topo.priority_order();
let mut stream_last: Vec<(u64, NodeId, StreamId)> = Vec::new();
let mut node_stream: HashMap<NodeId, StreamId> = HashMap::new();
for &node_id in &priority {
let preds = graph.predecessors(node_id)?;
let node_alap = topo.node_info(node_id).map(|i| i.alap).unwrap_or(0);
if preds.is_empty() {
let stream = if stream_last.len() < max_streams {
let sid = StreamId(stream_last.len() as u32);
stream_last.push((node_alap, node_id, sid));
sid
} else {
let best = stream_last
.iter()
.enumerate()
.min_by_key(|(_, (alap, _, _))| *alap)
.map(|(i, (_, _, sid))| (i, *sid))
.ok_or_else(|| {
GraphError::StreamPartitioningFailed(
"no streams available for assignment".into(),
)
})?;
stream_last[best.0] = (node_alap, node_id, best.1);
best.1
};
node_stream.insert(node_id, stream);
} else {
let pred_set: std::collections::HashSet<NodeId> = preds.iter().copied().collect();
let compatible_stream = stream_last
.iter()
.find(|(_, last_node, _)| pred_set.contains(last_node))
.map(|(_, _, sid)| *sid);
let stream = if let Some(sid) = compatible_stream {
sid
} else {
if stream_last.len() < max_streams {
let base_stream = preds
.iter()
.filter_map(|&p| node_stream.get(&p).copied())
.next()
.unwrap_or(StreamId::DEFAULT);
let base_last_is_pred = stream_last
.iter()
.find(|(_, _, sid)| *sid == base_stream)
.map(|(_, ln, _)| pred_set.contains(ln))
.unwrap_or(false);
if !base_last_is_pred {
let sid = StreamId(stream_last.len() as u32);
stream_last.push((node_alap, node_id, sid));
node_stream.insert(node_id, sid);
continue;
}
base_stream
} else {
preds
.iter()
.filter_map(|&p| node_stream.get(&p).copied())
.max_by_key(|&s| {
stream_last
.iter()
.find(|(_, _, sid)| *sid == s)
.map(|(alap, _, _)| *alap)
.unwrap_or(0)
})
.unwrap_or(StreamId::DEFAULT)
}
};
if let Some(entry) = stream_last.iter_mut().find(|(_, _, sid)| *sid == stream) {
if node_alap > entry.0 {
entry.0 = node_alap;
}
entry.1 = node_id;
}
node_stream.insert(node_id, stream);
}
}
let assignments: Vec<StreamAssignment> = topo
.order
.iter()
.map(|&nid| StreamAssignment {
node: nid,
stream: node_stream.get(&nid).copied().unwrap_or(StreamId::DEFAULT),
})
.collect();
let mut sync_points: Vec<SyncPoint> = Vec::new();
for edge in graph.edges() {
let (from, to) = edge;
let sf = node_stream.get(&from).copied().unwrap_or(StreamId::DEFAULT);
let st = node_stream.get(&to).copied().unwrap_or(StreamId::DEFAULT);
if sf != st {
sync_points.push(SyncPoint {
from,
stream_from: sf,
to,
stream_to: st,
});
}
}
sync_points.sort_by_key(|sp| (sp.from.0, sp.to.0));
sync_points.dedup_by_key(|sp| (sp.from, sp.to));
let num_streams = stream_last.len().max(1);
Ok(StreamPlan {
assignments,
sync_points,
num_streams,
max_streams,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::builder::GraphBuilder;
fn barrier(b: &mut GraphBuilder, name: &str) -> NodeId {
b.add_barrier(name)
}
#[test]
fn stream_empty_graph() {
let g = ComputeGraph::new();
assert!(matches!(analyse(&g, 4), Err(GraphError::EmptyGraph)));
}
#[test]
fn stream_single_node_default_stream() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let n = barrier(&mut b, "n");
let g = b.build().expect("test graph builds successfully");
let plan = analyse(&g, 4).expect("stream assignment analysis succeeds on valid graph");
assert_eq!(plan.stream_of(n), StreamId(0));
}
#[test]
fn stream_linear_chain_single_stream() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = barrier(&mut b, "a");
let bnode = barrier(&mut b, "b");
let c = barrier(&mut b, "c");
b.chain(&[a, bnode, c]);
let g = b.build().expect("test graph builds successfully");
let plan = analyse(&g, 4).expect("stream assignment analysis succeeds on valid graph");
assert!(!plan.has_concurrency());
assert_eq!(plan.sync_count(), 0);
}
#[test]
fn stream_fork_parallel_branches() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let src = barrier(&mut b, "src");
let a = barrier(&mut b, "a");
let bnode = barrier(&mut b, "b");
b.fan_out(src, &[a, bnode]);
let g = b.build().expect("test graph builds successfully");
let plan = analyse(&g, 4).expect("stream assignment analysis succeeds on valid graph");
let sa = plan.stream_of(a);
let sb = plan.stream_of(bnode);
assert_ne!(
sa, sb,
"independent branches should be on different streams"
);
}
#[test]
fn stream_max_streams_one_disables_parallelism() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = barrier(&mut b, "a");
let bnode = barrier(&mut b, "b");
let c = barrier(&mut b, "c");
let g = b.build().expect("test graph builds successfully");
let plan = analyse(&g, 1).expect("stream assignment analysis succeeds on valid graph");
assert_eq!(plan.stream_of(a), StreamId(0));
assert_eq!(plan.stream_of(bnode), StreamId(0));
assert_eq!(plan.stream_of(c), StreamId(0));
assert!(!plan.has_concurrency());
}
#[test]
fn stream_cross_stream_sync_detected() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = barrier(&mut b, "a");
let c = barrier(&mut b, "c");
let bnode = barrier(&mut b, "b");
b.dep(a, bnode).dep(c, bnode);
let g = b.build().expect("test graph builds successfully");
let plan = analyse(&g, 4).expect("stream assignment analysis succeeds on valid graph");
let sa = plan.stream_of(a);
let sc = plan.stream_of(c);
if sa != sc {
assert!(plan.sync_count() > 0);
}
}
#[test]
fn stream_nodes_on_stream_zero() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = barrier(&mut b, "a");
let bnode = barrier(&mut b, "b");
b.dep(a, bnode);
let g = b.build().expect("test graph builds successfully");
let plan = analyse(&g, 4).expect("stream assignment analysis succeeds on valid graph");
let on_s0 = plan.nodes_on(StreamId(0));
assert!(!on_s0.is_empty());
}
#[test]
fn stream_assignment_covers_all_nodes() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = barrier(&mut b, "a");
let bnode = barrier(&mut b, "b");
let c = barrier(&mut b, "c");
b.chain(&[a, bnode, c]);
let g = b.build().expect("test graph builds successfully");
let plan = analyse(&g, 2).expect("stream assignment analysis succeeds on valid graph");
assert_eq!(plan.assignments.len(), 3);
}
#[test]
fn stream_plan_respects_max_streams_cap() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
for i in 0..5 {
barrier(&mut b, &format!("n{i}"));
}
let g = b.build().expect("test graph builds successfully");
let plan = analyse(&g, 3).expect("stream assignment analysis succeeds on valid graph");
assert!(plan.num_streams <= 3);
}
#[test]
fn stream_diamond_join_handled() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = barrier(&mut b, "a");
let bnode = barrier(&mut b, "b");
let c = barrier(&mut b, "c");
let d = barrier(&mut b, "d");
b.dep(a, bnode).dep(a, c).dep(bnode, d).dep(c, d);
let g = b.build().expect("test graph builds successfully");
let plan = analyse(&g, 4).expect("stream assignment analysis succeeds on valid graph");
assert_eq!(plan.assignments.len(), 4);
}
#[test]
fn stream_zero_max_treated_as_one() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
barrier(&mut b, "n");
let g = b.build().expect("test graph builds successfully");
let plan =
analyse(&g, 0).expect("stream assignment analysis succeeds even with 0 max_streams"); assert_eq!(plan.num_streams, 1);
}
}