use std::collections::{HashMap, HashSet, VecDeque};
use crate::analysis::{dominance_analyse, topo_analyse};
use crate::error::{GraphError, GraphResult};
use crate::graph::ComputeGraph;
use crate::node::{BufferId, GraphNode, KernelConfig, NodeId, NodeKind};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ReductionPattern {
LayerNorm,
Softmax,
Generic,
}
impl ReductionPattern {
#[must_use]
pub fn name(self) -> &'static str {
match self {
Self::LayerNorm => "layernorm",
Self::Softmax => "softmax",
Self::Generic => "reduction",
}
}
}
impl std::fmt::Display for ReductionPattern {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.name())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ReductionFusionGroup {
pub id: usize,
pub root: NodeId,
pub sink: NodeId,
pub members: Vec<NodeId>,
pub pattern: ReductionPattern,
pub config: KernelConfig,
pub tag: String,
}
impl ReductionFusionGroup {
#[must_use]
pub fn size(&self) -> usize {
self.members.len()
}
#[must_use]
pub fn launches_saved(&self) -> usize {
self.members.len().saturating_sub(1)
}
}
#[derive(Debug, Clone, Default)]
pub struct ReductionFusionPlan {
pub groups: Vec<ReductionFusionGroup>,
pub node_to_group: HashMap<NodeId, usize>,
}
impl ReductionFusionPlan {
#[must_use]
pub fn fusion_count(&self) -> usize {
self.groups.len()
}
#[must_use]
pub fn nodes_saved(&self) -> usize {
self.groups
.iter()
.map(ReductionFusionGroup::launches_saved)
.sum()
}
#[must_use]
pub fn group_of(&self, node: NodeId) -> Option<&ReductionFusionGroup> {
self.node_to_group
.get(&node)
.and_then(|&idx| self.groups.get(idx))
}
#[must_use]
pub fn is_absorbed(&self, node: NodeId) -> bool {
match self.group_of(node) {
Some(g) => g.root != node,
None => false,
}
}
}
fn kernel_meta(graph: &ComputeGraph, node: NodeId) -> Option<(bool, KernelConfig)> {
match &graph.node(node).ok()?.kind {
NodeKind::KernelLaunch {
fusible, config, ..
} => Some((*fusible, *config)),
_ => None,
}
}
fn fn_name_lower(graph: &ComputeGraph, node: NodeId) -> String {
graph
.node(node)
.ok()
.and_then(|n| n.kind.function_name())
.unwrap_or("")
.to_ascii_lowercase()
}
fn configs_compatible(a: &KernelConfig, b: &KernelConfig) -> bool {
a.total_threads() == b.total_threads()
}
fn classify(graph: &ComputeGraph, members: &[NodeId]) -> ReductionPattern {
let names: Vec<String> = members.iter().map(|&m| fn_name_lower(graph, m)).collect();
let has = |needle: &str| names.iter().any(|n| n.contains(needle));
let softmax_like = has("exp") && (has("softmax") || has("div") || has("sum") || has("norm"));
let layernorm_like = (has("mean") || has("avg"))
&& (has("var") || has("std") || has("rms") || has("norm") || has("layernorm"));
if has("softmax") || softmax_like {
ReductionPattern::Softmax
} else if has("layernorm") || layernorm_like {
ReductionPattern::LayerNorm
} else {
ReductionPattern::Generic
}
}
fn grow_region(
graph: &ComputeGraph,
root: NodeId,
dt: &crate::analysis::DomTree,
topo_pos: &HashMap<NodeId, usize>,
claimed: &HashSet<NodeId>,
) -> GraphResult<Option<Vec<NodeId>>> {
let root_config = match kernel_meta(graph, root) {
Some((true, cfg)) => cfg,
_ => return Ok(None),
};
let member_ok = |n: NodeId| -> bool {
if n == root {
return true;
}
if claimed.contains(&n) {
return false;
}
if !dt.dominates(root, n) {
return false;
}
match kernel_meta(graph, n) {
Some((fusible, cfg)) => fusible && configs_compatible(&root_config, &cfg),
None => false,
}
};
let mut region: HashSet<NodeId> = HashSet::new();
region.insert(root);
let mut queue: VecDeque<NodeId> = VecDeque::new();
queue.push_back(root);
while let Some(cur) = queue.pop_front() {
for &succ in graph.successors(cur)? {
if region.contains(&succ) {
continue;
}
if member_ok(succ) {
region.insert(succ);
queue.push_back(succ);
}
}
}
if region.len() < 3 {
return Ok(None);
}
let mut exits: Vec<NodeId> = Vec::new();
for &m in ®ion {
let leaves = graph.successors(m)?.iter().any(|s| !region.contains(s));
let is_graph_sink = graph.successors(m)?.is_empty();
if leaves || is_graph_sink {
exits.push(m);
}
}
let sink = *region
.iter()
.max_by_key(|&&m| topo_pos.get(&m).copied().unwrap_or(0))
.ok_or_else(|| GraphError::Internal("reduction region unexpectedly empty".into()))?;
if !region_reaches_all(graph, ®ion, sink)? {
return Ok(None);
}
for &m in ®ion {
if m == sink {
continue;
}
let leaks = graph.successors(m)?.iter().any(|s| !region.contains(s));
if leaks {
return Ok(None);
}
}
let has_fanout = region
.iter()
.try_fold(false, |acc, &m| -> GraphResult<bool> {
if acc {
return Ok(true);
}
let in_region_succ = graph
.successors(m)?
.iter()
.filter(|s| region.contains(s))
.count();
Ok(in_region_succ >= 2)
})?;
if !has_fanout {
return Ok(None);
}
let mut members: Vec<NodeId> = region.into_iter().collect();
members.sort_by_key(|m| topo_pos.get(m).copied().unwrap_or(usize::MAX));
Ok(Some(members))
}
fn region_reaches_all(
graph: &ComputeGraph,
region: &HashSet<NodeId>,
sink: NodeId,
) -> GraphResult<bool> {
let mut reached: HashSet<NodeId> = HashSet::new();
reached.insert(sink);
let mut queue: VecDeque<NodeId> = VecDeque::new();
queue.push_back(sink);
while let Some(cur) = queue.pop_front() {
for &pred in graph.predecessors(cur)? {
if region.contains(&pred) && reached.insert(pred) {
queue.push_back(pred);
}
}
}
Ok(region.iter().all(|m| reached.contains(m)))
}
pub fn analyse(graph: &ComputeGraph) -> GraphResult<ReductionFusionPlan> {
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 claimed: HashSet<NodeId> = HashSet::new();
let mut groups: Vec<ReductionFusionGroup> = Vec::new();
let mut node_to_group: HashMap<NodeId, usize> = HashMap::new();
for &root in &topo.order {
if claimed.contains(&root) {
continue;
}
match kernel_meta(graph, root) {
Some((true, _)) => {}
_ => continue,
}
let members = match grow_region(graph, root, &dt, &topo_pos, &claimed)? {
Some(m) => m,
None => continue,
};
let sink = *members.last().ok_or_else(|| {
GraphError::Internal("reduction region members unexpectedly empty".into())
})?;
let pattern = classify(graph, &members);
let config = kernel_meta(graph, root)
.map(|(_, c)| c)
.unwrap_or_else(|| KernelConfig::linear(1, 1, 0));
let gid = groups.len();
let tag = format!(
"fused_{}_{}..{}",
pattern.name(),
graph.node(root)?.display_name(),
graph.node(sink)?.display_name()
);
for &m in &members {
claimed.insert(m);
node_to_group.insert(m, gid);
}
groups.push(ReductionFusionGroup {
id: gid,
root,
sink,
members,
pattern,
config,
tag,
});
}
Ok(ReductionFusionPlan {
groups,
node_to_group,
})
}
pub fn rewrite(graph: &ComputeGraph, plan: &ReductionFusionPlan) -> GraphResult<ComputeGraph> {
let mut out = ComputeGraph::new();
for buf in graph.buffers() {
out.add_buffer(buf.clone());
}
let mut old_to_new: HashMap<NodeId, NodeId> = HashMap::new();
for old in graph.nodes() {
let oid = old.id;
if let Some(group) = plan.group_of(oid) {
if group.root != oid {
continue;
}
let region: HashSet<NodeId> = group.members.iter().copied().collect();
let mut region_outputs: HashSet<BufferId> = HashSet::new();
for &m in &group.members {
for &b in &graph.node(m)?.outputs {
region_outputs.insert(b);
}
}
let mut fused_inputs: Vec<BufferId> = Vec::new();
let mut seen_in: HashSet<BufferId> = HashSet::new();
for &m in &group.members {
for &b in &graph.node(m)?.inputs {
if !region_outputs.contains(&b) && seen_in.insert(b) {
fused_inputs.push(b);
}
}
}
let fused_outputs: Vec<BufferId> = graph.node(group.sink)?.outputs.clone();
let fn_name = format!(
"{}_{}",
group.pattern.name(),
group
.members
.iter()
.filter_map(|&m| graph.node(m).ok().and_then(|n| n.kind.function_name()))
.collect::<Vec<_>>()
.join("_")
);
let cost: u64 = group
.members
.iter()
.filter_map(|&m| graph.node(m).ok().map(|n| n.cost_hint))
.sum();
let kind = NodeKind::KernelLaunch {
function_name: fn_name,
config: group.config,
fusible: true,
};
let node = GraphNode::new(NodeId(0), kind)
.with_inputs(fused_inputs)
.with_outputs(fused_outputs)
.with_cost(cost.max(1))
.with_name(group.tag.clone());
let nid = out.add_node(node);
for &m in ®ion {
old_to_new.insert(m, nid);
}
} else {
let mut node = GraphNode::new(NodeId(0), old.kind.clone())
.with_inputs(old.inputs.iter().copied())
.with_outputs(old.outputs.iter().copied())
.with_cost(old.cost_hint);
if let Some(s) = old.stream_hint {
node = node.with_stream(s);
}
if let Some(name) = &old.name {
node = node.with_name(name.clone());
}
let nid = out.add_node(node);
old_to_new.insert(oid, nid);
}
}
let mut added: HashSet<(NodeId, NodeId)> = HashSet::new();
for (from_old, to_old) in graph.edges() {
let from_new = *old_to_new
.get(&from_old)
.ok_or_else(|| GraphError::Internal("missing node mapping (from)".into()))?;
let to_new = *old_to_new
.get(&to_old)
.ok_or_else(|| GraphError::Internal("missing node mapping (to)".into()))?;
if from_new == to_new {
continue;
}
if added.insert((from_new, to_new)) {
out.add_edge(from_new, to_new)?;
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::builder::GraphBuilder;
use crate::executor::{ExecutionPlan, SequentialExecutor};
use crate::node::MemcpyDir;
fn build_layernorm() -> (ComputeGraph, Vec<NodeId>) {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let mean = b.add_kernel("mean", 4, 256, 0).fusible(true).finish();
let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
let var = b.add_kernel("variance", 4, 256, 0).fusible(true).finish();
let norm = b.add_kernel("normalize", 4, 256, 0).fusible(true).finish();
let scale = b
.add_kernel("scale_shift", 4, 256, 0)
.fusible(true)
.finish();
b.dep(mean, sub);
b.dep(sub, var);
b.dep(sub, norm); b.dep(var, norm);
b.dep(norm, scale);
let g = b.build().expect("layernorm graph builds");
(g, vec![mean, sub, var, norm, scale])
}
fn build_softmax() -> (ComputeGraph, Vec<NodeId>) {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
b.dep(mx, sub);
b.dep(sub, exp);
b.dep(exp, sum);
b.dep(exp, div); b.dep(sum, div);
let g = b.build().expect("softmax graph builds");
(g, vec![mx, sub, exp, sum, div])
}
#[test]
fn reduction_empty_graph_errors() {
let g = ComputeGraph::new();
assert!(matches!(analyse(&g), Err(GraphError::EmptyGraph)));
}
#[test]
fn layernorm_region_detected() {
let (g, ids) = build_layernorm();
let plan = analyse(&g).expect("reduction fusion analysis succeeds");
assert_eq!(plan.fusion_count(), 1);
let group = plan
.group_of(ids[0])
.expect("mean belongs to a fused region");
assert_eq!(group.size(), 5);
for id in &ids {
assert!(group.members.contains(id), "member {id} missing");
}
assert_eq!(group.root, ids[0]); assert_eq!(group.sink, ids[4]); assert_eq!(group.pattern, ReductionPattern::LayerNorm);
assert_eq!(plan.nodes_saved(), 4);
}
#[test]
fn softmax_region_detected() {
let (g, ids) = build_softmax();
let plan = analyse(&g).expect("reduction fusion analysis succeeds");
assert_eq!(plan.fusion_count(), 1);
let group = plan
.group_of(ids[2])
.expect("exp belongs to a fused region");
assert_eq!(group.size(), 5);
assert_eq!(group.root, ids[0]); assert_eq!(group.sink, ids[4]); assert_eq!(group.pattern, ReductionPattern::Softmax);
assert_eq!(plan.nodes_saved(), 4);
}
#[test]
fn absorbed_members_flagged() {
let (g, ids) = build_softmax();
let plan = analyse(&g).expect("reduction fusion analysis succeeds");
assert!(!plan.is_absorbed(ids[0]));
for id in &ids[1..] {
assert!(plan.is_absorbed(*id), "member {id} should be absorbed");
}
}
#[test]
fn linear_chain_not_a_reduction_region() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k0 = b.add_kernel("a", 4, 256, 0).fusible(true).finish();
let k1 = b.add_kernel("b", 4, 256, 0).fusible(true).finish();
let k2 = b.add_kernel("c", 4, 256, 0).fusible(true).finish();
b.chain(&[k0, k1, k2]);
let g = b.build().expect("chain graph builds");
let plan = analyse(&g).expect("reduction fusion analysis succeeds");
assert_eq!(plan.fusion_count(), 0);
assert_eq!(plan.nodes_saved(), 0);
}
#[test]
fn non_fusible_member_breaks_region() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
let exp = b.add_kernel("exp", 4, 256, 0).fusible(false).finish();
let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
b.dep(mx, sub);
b.dep(sub, exp);
b.dep(exp, sum);
b.dep(exp, div);
b.dep(sum, div);
let g = b.build().expect("graph builds");
let plan = analyse(&g).expect("reduction fusion analysis succeeds");
assert_eq!(plan.fusion_count(), 0);
}
#[test]
fn open_region_leaking_intermediate_not_fused() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
let leak = b.add_memcpy("leak", MemcpyDir::DeviceToHost, 1024);
b.dep(mx, sub);
b.dep(sub, exp);
b.dep(exp, sum);
b.dep(exp, div);
b.dep(sum, div);
b.dep(exp, leak); let g = b.build().expect("graph builds");
let plan = analyse(&g).expect("reduction fusion analysis succeeds");
assert_eq!(plan.fusion_count(), 0);
}
#[test]
fn incompatible_config_member_excluded() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
let div = b.add_kernel("divide", 8, 256, 0).fusible(true).finish(); b.dep(mx, sub);
b.dep(sub, exp);
b.dep(exp, sum);
b.dep(exp, div);
b.dep(sum, div);
let g = b.build().expect("graph builds");
let plan = analyse(&g).expect("reduction fusion analysis succeeds");
assert_eq!(plan.fusion_count(), 0);
}
#[test]
fn rewrite_collapses_region_to_one_node() {
let (g, _ids) = build_layernorm();
let plan = analyse(&g).expect("analysis succeeds");
let fused = rewrite(&g, &plan).expect("rewrite succeeds");
assert_eq!(g.node_count(), 5);
assert_eq!(fused.node_count(), 1);
let only = fused.node(NodeId(0)).expect("fused node exists");
assert!(only.kind.is_compute());
assert!(
only.kind
.function_name()
.unwrap_or("")
.contains("layernorm")
);
}
#[test]
fn rewrite_preserves_external_topology() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let up = b.add_memcpy("up", MemcpyDir::HostToDevice, 1024);
let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
let dn = b.add_memcpy("dn", MemcpyDir::DeviceToHost, 1024);
b.dep(up, mx);
b.dep(mx, sub);
b.dep(sub, exp);
b.dep(exp, sum);
b.dep(exp, div);
b.dep(sum, div);
b.dep(div, dn);
let g = b.build().expect("graph builds");
let plan = analyse(&g).expect("analysis succeeds");
assert_eq!(plan.fusion_count(), 1);
let fused = rewrite(&g, &plan).expect("rewrite succeeds");
assert_eq!(fused.node_count(), 3);
let up_new = fused.sources();
assert_eq!(up_new.len(), 1);
let dn_new = fused.sinks();
assert_eq!(dn_new.len(), 1);
assert!(fused.is_reachable(up_new[0], dn_new[0]));
assert_eq!(fused.kernel_nodes().len(), 1);
}
#[test]
fn rewrite_no_match_is_identity() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let k0 = b.add_kernel("a", 4, 256, 0).fusible(true).finish();
let k1 = b.add_kernel("b", 4, 256, 0).fusible(true).finish();
b.chain(&[k0, k1]);
let g = b.build().expect("graph builds");
let plan = analyse(&g).expect("analysis succeeds");
let fused = rewrite(&g, &plan).expect("rewrite succeeds");
assert_eq!(fused.node_count(), g.node_count());
assert_eq!(fused.edge_count(), g.edge_count());
}
#[test]
fn simulator_agrees_before_and_after_fusion() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let up = b.add_memcpy("up", MemcpyDir::HostToDevice, 4096);
let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
let dn = b.add_memcpy("dn", MemcpyDir::DeviceToHost, 4096);
b.dep(up, mx);
b.dep(mx, sub);
b.dep(sub, exp);
b.dep(exp, sum);
b.dep(exp, div);
b.dep(sum, div);
b.dep(div, dn);
let g = b.build().expect("graph builds");
let plan = analyse(&g).expect("analysis succeeds");
assert_eq!(plan.fusion_count(), 1);
let fused = rewrite(&g, &plan).expect("rewrite succeeds");
let before =
SequentialExecutor::new(&ExecutionPlan::build(&g, 4).expect("plan(before) builds"))
.run()
.expect("before runs");
let after =
SequentialExecutor::new(&ExecutionPlan::build(&fused, 4).expect("plan(after) builds"))
.run()
.expect("after runs");
assert_eq!(before.bytes_copied, after.bytes_copied);
assert_eq!(before.bytes_copied, 4096 * 2);
assert_eq!(before.bytes_set, after.bytes_set);
assert_eq!(after.kernels_launched, 1);
assert!(after.kernels_launched <= before.kernels_launched);
}
#[test]
fn pattern_display_and_name() {
assert_eq!(ReductionPattern::LayerNorm.name(), "layernorm");
assert_eq!(ReductionPattern::Softmax.to_string(), "softmax");
assert_eq!(ReductionPattern::Generic.name(), "reduction");
}
#[test]
fn generic_reduction_region_classified() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let r = b.add_kernel("reduce_op", 4, 256, 0).fusible(true).finish();
let a = b.add_kernel("elemwise_a", 4, 256, 0).fusible(true).finish();
let c = b.add_kernel("elemwise_c", 4, 256, 0).fusible(true).finish();
let join = b.add_kernel("combine", 4, 256, 0).fusible(true).finish();
b.dep(r, a);
b.dep(r, c);
b.dep(a, join);
b.dep(c, join);
let g = b.build().expect("graph builds");
let plan = analyse(&g).expect("analysis succeeds");
assert_eq!(plan.fusion_count(), 1);
let group = plan.group_of(r).expect("r in a region");
assert_eq!(group.pattern, ReductionPattern::Generic);
assert_eq!(group.size(), 4);
}
}