use std::collections::{HashMap, HashSet};
use crate::analysis::{dominance_analyse, topo_analyse};
use crate::error::{GraphError, GraphResult};
use crate::graph::ComputeGraph;
use crate::node::{KernelConfig, NodeId, NodeKind};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FusionGroup {
pub id: usize,
pub members: Vec<NodeId>,
pub config: KernelConfig,
pub tag: String,
}
impl FusionGroup {
#[must_use]
pub fn size(&self) -> usize {
self.members.len()
}
#[must_use]
pub fn is_trivial(&self) -> bool {
self.members.len() == 1
}
}
#[derive(Debug, Clone)]
pub struct FusionPlan {
pub groups: Vec<FusionGroup>,
pub node_to_group: HashMap<NodeId, usize>,
}
impl FusionPlan {
pub fn fusion_count(&self) -> usize {
self.groups.iter().filter(|g| !g.is_trivial()).count()
}
pub fn nodes_saved(&self) -> usize {
self.groups
.iter()
.filter(|g| !g.is_trivial())
.map(|g| g.size() - 1)
.sum()
}
pub fn group_of(&self, node: NodeId) -> Option<&FusionGroup> {
self.node_to_group
.get(&node)
.and_then(|&idx| self.groups.get(idx))
}
}
fn configs_compatible(a: &KernelConfig, b: &KernelConfig) -> bool {
a.total_threads() == b.total_threads()
}
fn only_fusible_between(
graph: &ComputeGraph,
a: NodeId,
b: NodeId,
topo_pos: &HashMap<NodeId, usize>,
) -> bool {
let pos_a = topo_pos[&a];
let pos_b = topo_pos[&b];
if pos_b <= pos_a + 1 {
return true; }
let mut visited = HashSet::new();
let mut stack = vec![a];
while let Some(cur) = stack.pop() {
if cur == b {
continue;
}
for &s in graph.successors(cur).unwrap_or(&[]) {
if visited.insert(s) {
if s == b {
continue;
}
let node = graph.node(s).ok();
let is_fusible = node.map(|n| n.kind.is_fusible()).unwrap_or(false);
let is_barrier = node
.map(|n| matches!(n.kind, NodeKind::Barrier))
.unwrap_or(false);
let spos = topo_pos.get(&s).copied().unwrap_or(usize::MAX);
if spos < pos_b && (is_fusible || is_barrier) {
stack.push(s);
} else if spos < pos_b && !is_fusible && !is_barrier {
return false; }
}
}
}
true
}
pub fn analyse(graph: &ComputeGraph) -> GraphResult<FusionPlan> {
if graph.is_empty() {
return Err(GraphError::EmptyGraph);
}
let topo = topo_analyse(graph)?;
let dt = dominance_analyse(graph)?;
let topo_pos: HashMap<NodeId, usize> = topo
.order
.iter()
.enumerate()
.map(|(p, &id)| (id, p))
.collect();
let mut assigned: HashMap<NodeId, usize> = HashMap::new();
let mut groups: Vec<FusionGroup> = Vec::new();
for &node_id in &topo.order {
if assigned.contains_key(&node_id) {
continue;
}
let node = graph.node(node_id)?;
let (is_fusible, base_config) = match &node.kind {
NodeKind::KernelLaunch {
fusible, config, ..
} => (*fusible, *config),
_ => {
let gid = groups.len();
groups.push(FusionGroup {
id: gid,
members: vec![node_id],
config: KernelConfig::linear(1, 1, 0),
tag: format!("non_kernel_{}", node.kind.tag()),
});
assigned.insert(node_id, gid);
continue;
}
};
if !is_fusible {
let gid = groups.len();
groups.push(FusionGroup {
id: gid,
members: vec![node_id],
config: base_config,
tag: format!("non_fusible_{}", node.display_name()),
});
assigned.insert(node_id, gid);
continue;
}
let gid = groups.len();
let mut members = vec![node_id];
assigned.insert(node_id, gid);
let mut frontier = graph.successors(node_id)?.to_vec();
while let Some(succ_id) = frontier.first().copied() {
frontier.remove(0);
if assigned.contains_key(&succ_id) {
continue;
}
let succ = graph.node(succ_id)?;
let (succ_fusible, succ_config) = match &succ.kind {
NodeKind::KernelLaunch {
fusible, config, ..
} => (*fusible, *config),
_ => continue,
};
if !succ_fusible {
continue;
}
if !configs_compatible(&base_config, &succ_config) {
continue;
}
let last_member = *members.last().ok_or_else(|| {
GraphError::Internal("fusion group members unexpectedly empty".into())
})?;
if !dt.dominates(last_member, succ_id) {
continue;
}
if !only_fusible_between(graph, last_member, succ_id, &topo_pos) {
continue;
}
members.push(succ_id);
assigned.insert(succ_id, gid);
for &next in graph.successors(succ_id)? {
if !assigned.contains_key(&next) {
frontier.push(next);
}
}
}
let tag = if members.len() > 1 {
format!(
"fused_{}..{}",
graph.node(members[0])?.display_name(),
graph
.node(*members.last().ok_or_else(|| {
GraphError::Internal("fusion group members unexpectedly empty".into())
})?)?
.display_name()
)
} else {
format!("solo_{}", graph.node(node_id)?.display_name())
};
groups.push(FusionGroup {
id: gid,
members,
config: base_config,
tag,
});
}
let node_to_group: HashMap<NodeId, usize> = assigned;
Ok(FusionPlan {
groups,
node_to_group,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::builder::GraphBuilder;
use crate::node::MemcpyDir;
fn fusible_kernel(b: &mut GraphBuilder, name: &str) -> NodeId {
b.add_kernel(name, 4, 256, 0).fusible(true).finish()
}
fn non_fusible_kernel(b: &mut GraphBuilder, name: &str) -> NodeId {
b.add_kernel(name, 4, 256, 0).fusible(false).finish()
}
#[test]
fn fusion_empty_graph() {
let g = ComputeGraph::new();
assert!(matches!(analyse(&g), Err(GraphError::EmptyGraph)));
}
#[test]
fn fusion_single_fusible_kernel_trivial_group() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k = fusible_kernel(&mut b, "add");
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
assert_eq!(plan.groups.len(), 1);
assert!(plan.groups[0].is_trivial());
assert_eq!(plan.group_of(k).unwrap().members, vec![k]);
}
#[test]
fn fusion_chain_of_fusible_kernels_merged() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k0 = fusible_kernel(&mut b, "k0");
let k1 = fusible_kernel(&mut b, "k1");
let k2 = fusible_kernel(&mut b, "k2");
b.chain(&[k0, k1, k2]);
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
assert_eq!(plan.fusion_count(), 1);
let group = plan.group_of(k0).unwrap();
assert_eq!(group.size(), 3);
assert!(group.members.contains(&k0));
assert!(group.members.contains(&k1));
assert!(group.members.contains(&k2));
}
#[test]
fn fusion_non_fusible_breaks_chain() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k0 = fusible_kernel(&mut b, "k0");
let k1 = non_fusible_kernel(&mut b, "k1");
let k2 = fusible_kernel(&mut b, "k2");
b.chain(&[k0, k1, k2]);
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
let g0 = plan.group_of(k0).unwrap().id;
let g2 = plan.group_of(k2).unwrap().id;
assert_ne!(g0, g2);
}
#[test]
fn fusion_memcpy_not_fused() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let upload = b.add_memcpy("up", MemcpyDir::HostToDevice, 1024);
let k = fusible_kernel(&mut b, "k");
let download = b.add_memcpy("dn", MemcpyDir::DeviceToHost, 1024);
b.chain(&[upload, k, download]);
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
let gup = plan.group_of(upload).unwrap();
let gdn = plan.group_of(download).unwrap();
assert!(gup.is_trivial());
assert!(gdn.is_trivial());
}
#[test]
fn fusion_incompatible_configs_not_fused() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k0 = b.add_kernel("k0", 4, 256, 0).fusible(true).finish();
let k1 = b.add_kernel("k1", 8, 256, 0).fusible(true).finish();
b.dep(k0, k1);
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
let gk0 = plan.group_of(k0).unwrap().id;
let gk1 = plan.group_of(k1).unwrap().id;
assert_ne!(gk0, gk1);
}
#[test]
fn fusion_nodes_saved_count() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k0 = fusible_kernel(&mut b, "k0");
let k1 = fusible_kernel(&mut b, "k1");
let k2 = fusible_kernel(&mut b, "k2");
b.chain(&[k0, k1, k2]);
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
assert_eq!(plan.nodes_saved(), 2);
}
#[test]
fn fusion_plan_covers_all_nodes() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k0 = fusible_kernel(&mut b, "k0");
let k1 = non_fusible_kernel(&mut b, "k1");
let upload = b.add_memcpy("up", MemcpyDir::HostToDevice, 512);
b.chain(&[upload, k0, k1]);
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
let total: usize = plan.groups.iter().map(|g| g.size()).sum();
assert_eq!(total, 3);
assert!(plan.node_to_group.contains_key(&k0));
assert!(plan.node_to_group.contains_key(&k1));
assert!(plan.node_to_group.contains_key(&upload));
}
#[test]
fn fusion_parallel_branches_not_fused() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let src = b.add_barrier("src");
let k0 = fusible_kernel(&mut b, "k0");
let k1 = fusible_kernel(&mut b, "k1");
b.fan_out(src, &[k0, k1]);
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
let gk0 = plan.group_of(k0).unwrap().id;
let gk1 = plan.group_of(k1).unwrap().id;
assert_ne!(gk0, gk1);
}
#[test]
fn fusion_group_tag_contains_names() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k0 = fusible_kernel(&mut b, "relu");
let k1 = fusible_kernel(&mut b, "scale");
b.dep(k0, k1);
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
let group = plan.group_of(k0).unwrap();
assert!(!group.tag.is_empty());
}
#[test]
fn fusion_empty_fusible_graph_one_group() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k = fusible_kernel(&mut b, "solo");
let g = b.build().unwrap();
let plan = analyse(&g).unwrap();
assert_eq!(plan.fusion_count(), 0); assert_eq!(plan.nodes_saved(), 0);
assert_eq!(plan.group_of(k).unwrap().size(), 1);
}
}