use std::collections::{HashMap, HashSet, VecDeque};
use crate::GpuOptimError;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum OpKind {
Add,
Mul,
Sub,
Div,
Relu,
Sigmoid,
Tanh,
Scale,
Axpy,
MatMul,
Reduce,
Transpose,
}
impl OpKind {
pub fn is_fusible(&self) -> bool {
matches!(
self,
OpKind::Add
| OpKind::Mul
| OpKind::Sub
| OpKind::Div
| OpKind::Relu
| OpKind::Sigmoid
| OpKind::Tanh
| OpKind::Scale
| OpKind::Axpy
)
}
pub fn is_barrier(&self) -> bool {
!self.is_fusible()
}
}
#[derive(Debug, Clone)]
pub struct FusionOp {
pub id: usize,
pub kind: OpKind,
pub inputs: Vec<usize>,
pub output_shape: Vec<usize>,
pub dtype_bytes: usize,
}
impl FusionOp {
pub fn output_elements(&self) -> Result<u64, GpuOptimError> {
let mut elements: u64 = 1;
for &dim in &self.output_shape {
elements = elements.checked_mul(dim as u64).ok_or_else(|| {
GpuOptimError::UnsupportedOperation(format!(
"op {} output shape {:?} overflows the element counter",
self.id, self.output_shape
))
})?;
}
Ok(elements)
}
pub fn output_bytes(&self) -> Result<u64, GpuOptimError> {
let elements = self.output_elements()?;
elements
.checked_mul(self.dtype_bytes as u64)
.ok_or_else(|| {
GpuOptimError::UnsupportedOperation(format!(
"op {} output ({} elements x {} bytes) overflows the byte counter",
self.id, elements, self.dtype_bytes
))
})
}
}
#[derive(Debug, Default, Clone)]
pub struct FusionGraph {
ops: Vec<FusionOp>,
}
impl FusionGraph {
pub fn new() -> Self {
Self { ops: Vec::new() }
}
pub fn add_op(
&mut self,
kind: OpKind,
inputs: Vec<usize>,
output_shape: Vec<usize>,
dtype_bytes: usize,
) -> usize {
let id = self.ops.len();
self.ops.push(FusionOp {
id,
kind,
inputs,
output_shape,
dtype_bytes,
});
id
}
pub fn ops(&self) -> &[FusionOp] {
&self.ops
}
pub fn num_ops(&self) -> usize {
self.ops.len()
}
pub fn is_empty(&self) -> bool {
self.ops.is_empty()
}
pub fn validate(&self) -> Result<(), GpuOptimError> {
let n = self.ops.len();
for op in &self.ops {
if op.dtype_bytes == 0 {
return Err(GpuOptimError::InvalidState(format!(
"op {} has dtype_bytes == 0",
op.id
)));
}
for &producer in &op.inputs {
if producer >= n {
return Err(GpuOptimError::InvalidState(format!(
"op {} references non-existent input op {}",
op.id, producer
)));
}
if producer == op.id {
return Err(GpuOptimError::InvalidState(format!(
"op {} references itself as an input",
op.id
)));
}
}
op.output_bytes()?;
}
self.topological_order()?;
Ok(())
}
fn topological_order(&self) -> Result<Vec<usize>, GpuOptimError> {
let n = self.ops.len();
let mut indegree = vec![0usize; n];
let mut adjacency: Vec<Vec<usize>> = vec![Vec::new(); n];
for consumer in &self.ops {
let mut seen: HashSet<usize> = HashSet::new();
for &producer in &consumer.inputs {
if producer >= n {
return Err(GpuOptimError::InvalidState(format!(
"op {} references non-existent input op {}",
consumer.id, producer
)));
}
if !seen.insert(producer) {
continue;
}
adjacency[producer].push(consumer.id);
indegree[consumer.id] += 1;
}
}
let mut queue: VecDeque<usize> = (0..n).filter(|&i| indegree[i] == 0).collect();
let mut order = Vec::with_capacity(n);
while let Some(node) = queue.pop_front() {
order.push(node);
for &consumer in &adjacency[node] {
indegree[consumer] -= 1;
if indegree[consumer] == 0 {
queue.push_back(consumer);
}
}
}
if order.len() != n {
return Err(GpuOptimError::InvalidState(
"operation graph contains a cycle".to_string(),
));
}
Ok(order)
}
}
#[derive(Debug, Clone)]
pub struct FusionGroup {
pub members: Vec<usize>,
pub external_inputs: Vec<usize>,
pub external_outputs: Vec<usize>,
pub internal_intermediates: Vec<usize>,
pub bytes_read: u64,
pub bytes_written: u64,
}
impl FusionGroup {
pub fn bytes_fused(&self) -> u64 {
self.bytes_read + self.bytes_written
}
pub fn is_fused(&self) -> bool {
self.members.len() > 1
}
}
#[derive(Debug, Clone)]
pub struct FusionPlan {
pub groups: Vec<FusionGroup>,
pub bytes_unfused: u64,
pub bytes_fused: u64,
pub bytes_saved: u64,
pub speedup_estimate: f64,
}
impl FusionPlan {
pub fn num_groups(&self) -> usize {
self.groups.len()
}
}
#[derive(Debug, Clone)]
pub struct FusionPlanner {
allow_broadcast: bool,
}
impl Default for FusionPlanner {
fn default() -> Self {
Self::new()
}
}
impl FusionPlanner {
pub fn new() -> Self {
Self {
allow_broadcast: true,
}
}
pub fn with_broadcast(mut self, allow_broadcast: bool) -> Self {
self.allow_broadcast = allow_broadcast;
self
}
fn shapes_fuse_compatible(&self, producer: &[usize], consumer: &[usize]) -> bool {
if producer == consumer {
return true;
}
if !self.allow_broadcast {
return false;
}
if producer.len() > consumer.len() {
return false;
}
let offset = consumer.len() - producer.len();
for (i, &producer_dim) in producer.iter().enumerate() {
let consumer_dim = consumer[offset + i];
if producer_dim != consumer_dim && producer_dim != 1 {
return false;
}
}
true
}
pub fn plan(&self, graph: &FusionGraph) -> Result<FusionPlan, GpuOptimError> {
graph.validate()?;
let ops = graph.ops();
let n = ops.len();
let mut bytes: Vec<u64> = Vec::with_capacity(n);
for op in ops {
bytes.push(op.output_bytes()?);
}
let topo = graph.topological_order()?;
let mut consumers: Vec<Vec<usize>> = vec![Vec::new(); n];
for consumer in ops {
let mut seen: HashSet<usize> = HashSet::new();
for &producer in &consumer.inputs {
if seen.insert(producer) {
consumers[producer].push(consumer.id);
}
}
}
let mut group_of: Vec<usize> = (0..n).collect();
for &consumer_id in &topo {
let consumer = &ops[consumer_id];
if !consumer.kind.is_fusible() {
continue;
}
let mut seen: HashSet<usize> = HashSet::new();
for &producer_id in &consumer.inputs {
if !seen.insert(producer_id) {
continue;
}
let producer = &ops[producer_id];
if !producer.kind.is_fusible() {
continue;
}
if !self.shapes_fuse_compatible(&producer.output_shape, &consumer.output_shape) {
continue;
}
let group_producer = group_of[producer_id];
let group_consumer = group_of[consumer_id];
if group_producer == group_consumer {
continue;
}
if merge_keeps_acyclic(ops, &group_of, group_producer, group_consumer) {
for label in group_of.iter_mut() {
if *label == group_consumer {
*label = group_producer;
}
}
}
}
}
let mut label_to_members: HashMap<usize, Vec<usize>> = HashMap::new();
for &id in &topo {
label_to_members.entry(group_of[id]).or_default().push(id);
}
let mut raw_groups: Vec<Vec<usize>> = label_to_members.into_values().collect();
for members in raw_groups.iter_mut() {
members.sort_unstable();
}
raw_groups.sort_by_key(|members| members[0]);
let mut groups: Vec<FusionGroup> = Vec::with_capacity(raw_groups.len());
let mut bytes_fused: u64 = 0;
for members in raw_groups {
let member_set: HashSet<usize> = members.iter().copied().collect();
let mut external_inputs: Vec<usize> = Vec::new();
let mut external_input_seen: HashSet<usize> = HashSet::new();
for &member in &members {
let mut seen: HashSet<usize> = HashSet::new();
for &producer in &ops[member].inputs {
if !seen.insert(producer) {
continue;
}
if !member_set.contains(&producer) && external_input_seen.insert(producer) {
external_inputs.push(producer);
}
}
}
external_inputs.sort_unstable();
let mut external_outputs: Vec<usize> = Vec::new();
let mut internal_intermediates: Vec<usize> = Vec::new();
for &member in &members {
let consumed_externally = consumers[member].iter().any(|c| !member_set.contains(c));
let is_terminal = consumers[member].is_empty();
if consumed_externally || is_terminal {
external_outputs.push(member);
} else {
internal_intermediates.push(member);
}
}
let bytes_read: u64 = external_inputs.iter().map(|&p| bytes[p]).sum();
let bytes_written: u64 = external_outputs.iter().map(|&m| bytes[m]).sum();
bytes_fused += bytes_read + bytes_written;
groups.push(FusionGroup {
members,
external_inputs,
external_outputs,
internal_intermediates,
bytes_read,
bytes_written,
});
}
let mut bytes_unfused: u64 = 0;
for op in ops {
let mut seen: HashSet<usize> = HashSet::new();
let mut read: u64 = 0;
for &producer in &op.inputs {
if seen.insert(producer) {
read += bytes[producer];
}
}
bytes_unfused += read + bytes[op.id];
}
let bytes_saved = bytes_unfused.saturating_sub(bytes_fused);
let speedup_estimate = if bytes_fused == 0 {
1.0
} else {
bytes_unfused as f64 / bytes_fused as f64
};
Ok(FusionPlan {
groups,
bytes_unfused,
bytes_fused,
bytes_saved,
speedup_estimate,
})
}
}
fn merge_keeps_acyclic(
ops: &[FusionOp],
group_of: &[usize],
group_a: usize,
group_b: usize,
) -> bool {
let label = |op_id: usize| -> usize {
let group = group_of[op_id];
if group == group_b {
group_a
} else {
group
}
};
let mut adjacency: HashMap<usize, HashSet<usize>> = HashMap::new();
let mut nodes: HashSet<usize> = HashSet::new();
for consumer in ops {
let consumer_label = label(consumer.id);
nodes.insert(consumer_label);
for &producer in &consumer.inputs {
let producer_label = label(producer);
nodes.insert(producer_label);
if producer_label != consumer_label {
adjacency
.entry(producer_label)
.or_default()
.insert(consumer_label);
}
}
}
let mut indegree: HashMap<usize, usize> = nodes.iter().map(|&node| (node, 0usize)).collect();
for targets in adjacency.values() {
for &target in targets {
if let Some(degree) = indegree.get_mut(&target) {
*degree += 1;
}
}
}
let mut queue: VecDeque<usize> = indegree
.iter()
.filter_map(|(&node, °ree)| if degree == 0 { Some(node) } else { None })
.collect();
let mut visited = 0usize;
while let Some(node) = queue.pop_front() {
visited += 1;
if let Some(targets) = adjacency.get(&node) {
for &target in targets {
if let Some(degree) = indegree.get_mut(&target) {
*degree -= 1;
if *degree == 0 {
queue.push_back(target);
}
}
}
}
}
visited == nodes.len()
}
#[cfg(test)]
mod tests {
use super::*;
fn find_group(plan: &FusionPlan, op_id: usize) -> &FusionGroup {
plan.groups
.iter()
.find(|g| g.members.contains(&op_id))
.expect("every op must belong to exactly one group")
}
#[test]
fn is_fusible_classification() {
assert!(OpKind::Add.is_fusible());
assert!(OpKind::Mul.is_fusible());
assert!(OpKind::Axpy.is_fusible());
assert!(OpKind::Scale.is_fusible());
assert!(!OpKind::MatMul.is_fusible());
assert!(!OpKind::Reduce.is_fusible());
assert!(!OpKind::Transpose.is_fusible());
assert!(OpKind::MatMul.is_barrier());
assert!(!OpKind::Relu.is_barrier());
}
#[test]
fn linear_chain_fuses_into_one_group() {
let mut graph = FusionGraph::new();
let a = graph.add_op(OpKind::Relu, vec![], vec![256], 4);
let b = graph.add_op(OpKind::Sigmoid, vec![a], vec![256], 4);
let c = graph.add_op(OpKind::Tanh, vec![b], vec![256], 4);
let plan = FusionPlanner::new()
.plan(&graph)
.expect("fusible chain must plan");
assert_eq!(plan.groups.len(), 1);
let group = &plan.groups[0];
assert_eq!(group.members, vec![a, b, c]);
assert!(group.external_inputs.is_empty());
assert_eq!(group.external_outputs, vec![c]);
assert_eq!(group.internal_intermediates, vec![a, b]);
let tensor = 256u64 * 4;
assert_eq!(plan.bytes_unfused, 5 * tensor);
assert_eq!(plan.bytes_fused, tensor);
assert!(plan.bytes_fused < plan.bytes_unfused);
assert_eq!(plan.bytes_saved, 4 * tensor);
assert!((plan.speedup_estimate - 5.0).abs() < 1e-9);
}
#[test]
fn barrier_splits_into_three_groups() {
let mut graph = FusionGraph::new();
let a = graph.add_op(OpKind::Relu, vec![], vec![128], 4);
let b = graph.add_op(OpKind::Sigmoid, vec![a], vec![128], 4);
let c = graph.add_op(OpKind::MatMul, vec![b], vec![128], 4); let d = graph.add_op(OpKind::Relu, vec![c], vec![128], 4);
let e = graph.add_op(OpKind::Tanh, vec![d], vec![128], 4);
let plan = FusionPlanner::new()
.plan(&graph)
.expect("graph with a barrier must plan");
assert_eq!(plan.groups.len(), 3);
assert_eq!(find_group(&plan, a).members, vec![a, b]);
assert_eq!(find_group(&plan, c).members, vec![c]);
assert_eq!(find_group(&plan, d).members, vec![d, e]);
}
#[test]
fn external_consumer_materializes_intermediate() {
let mut graph = FusionGraph::new();
let a = graph.add_op(OpKind::Relu, vec![], vec![256], 4);
let b = graph.add_op(OpKind::Sigmoid, vec![a], vec![256], 4);
let c = graph.add_op(OpKind::Tanh, vec![b], vec![256], 4);
let d = graph.add_op(OpKind::MatMul, vec![b], vec![256], 4);
let plan = FusionPlanner::new()
.plan(&graph)
.expect("diamond graph must plan");
assert_eq!(plan.groups.len(), 2);
let group = find_group(&plan, b);
assert_eq!(group.members, vec![a, b, c]);
assert!(
group.external_outputs.contains(&b),
"b is consumed outside the group and must be materialized"
);
assert!(!group.internal_intermediates.contains(&b));
assert!(group.external_outputs.contains(&c));
assert_eq!(group.internal_intermediates, vec![a]);
assert_eq!(find_group(&plan, d).members, vec![d]);
let tensor = 256u64 * 4;
assert_eq!(plan.bytes_unfused, 7 * tensor);
assert_eq!(plan.bytes_fused, 4 * tensor);
assert_eq!(plan.bytes_saved, 3 * tensor);
assert!((plan.speedup_estimate - 1.75).abs() < 1e-9);
}
#[test]
fn shared_external_input_counted_once() {
let mut graph = FusionGraph::new();
let x = graph.add_op(OpKind::MatMul, vec![], vec![16], 4);
let r = graph.add_op(OpKind::Relu, vec![x], vec![16], 4);
let s = graph.add_op(OpKind::Add, vec![r, x], vec![16], 4);
let plan = FusionPlanner::new().plan(&graph).expect("graph must plan");
assert_eq!(plan.groups.len(), 2);
let group = find_group(&plan, r);
assert_eq!(group.members, vec![r, s]);
assert_eq!(group.external_inputs, vec![x]);
let tensor = 16u64 * 4; assert_eq!(plan.bytes_unfused, 6 * tensor);
assert_eq!(plan.bytes_fused, 3 * tensor);
assert_eq!(plan.bytes_saved, 3 * tensor);
assert!((plan.speedup_estimate - 2.0).abs() < 1e-9);
}
#[test]
fn hand_computed_bytes_exact() {
let mut graph = FusionGraph::new();
let a = graph.add_op(OpKind::Mul, vec![], vec![10], 4); let b = graph.add_op(OpKind::Add, vec![a], vec![10], 4);
let plan = FusionPlanner::new()
.plan(&graph)
.expect("two-op chain must plan");
assert_eq!(plan.groups.len(), 1);
let group = &plan.groups[0];
assert!(group.external_inputs.is_empty());
assert_eq!(group.external_outputs, vec![b]);
assert_eq!(group.internal_intermediates, vec![a]);
assert_eq!(group.bytes_read, 0);
assert_eq!(group.bytes_written, 40);
assert_eq!(plan.bytes_unfused, 120);
assert_eq!(plan.bytes_fused, 40);
assert_eq!(plan.bytes_saved, 80);
assert!((plan.speedup_estimate - 3.0).abs() < 1e-9);
}
#[test]
fn fusion_avoids_introducing_cycle() {
let mut graph = FusionGraph::new();
let a = graph.add_op(OpKind::Relu, vec![], vec![16], 4);
let b = graph.add_op(OpKind::Sigmoid, vec![a], vec![16], 4);
let c = graph.add_op(OpKind::MatMul, vec![b], vec![16], 4); let d = graph.add_op(OpKind::Add, vec![a, c], vec![16], 4);
let plan = FusionPlanner::new().plan(&graph).expect("graph must plan");
assert_eq!(plan.groups.len(), 3);
assert_eq!(find_group(&plan, a).members, vec![a, b]);
assert_eq!(find_group(&plan, c).members, vec![c]);
assert_eq!(find_group(&plan, d).members, vec![d]);
assert!(find_group(&plan, a).external_outputs.contains(&a));
}
#[test]
fn broadcast_compatible_edge_fuses() {
let mut graph = FusionGraph::new();
let a = graph.add_op(OpKind::Relu, vec![], vec![1], 4); let b = graph.add_op(OpKind::Add, vec![a], vec![32], 4);
let plan = FusionPlanner::new()
.plan(&graph)
.expect("broadcast chain must plan");
assert_eq!(plan.groups.len(), 1);
assert_eq!(plan.groups[0].members, vec![a, b]);
let strict = FusionPlanner::new()
.with_broadcast(false)
.plan(&graph)
.expect("strict planner must plan");
assert_eq!(strict.groups.len(), 2);
}
#[test]
fn bytes_saved_and_speedup_invariants() {
let mut graph = FusionGraph::new();
let a = graph.add_op(OpKind::Scale, vec![], vec![64, 64], 4);
let b = graph.add_op(OpKind::Relu, vec![a], vec![64, 64], 4);
let c = graph.add_op(OpKind::Sigmoid, vec![b], vec![64, 64], 4);
graph.add_op(OpKind::Tanh, vec![c], vec![64, 64], 4);
let plan = FusionPlanner::new()
.plan(&graph)
.expect("fusible chain must plan");
assert_eq!(plan.groups.len(), 1);
assert_eq!(plan.bytes_saved, plan.bytes_unfused - plan.bytes_fused);
assert!(plan.speedup_estimate >= 1.0);
assert!(plan.bytes_fused < plan.bytes_unfused);
}
#[test]
fn cyclic_graph_rejected() {
let mut graph = FusionGraph::new();
let _a = graph.add_op(OpKind::Add, vec![1], vec![8], 4); let _b = graph.add_op(OpKind::Add, vec![0], vec![8], 4);
let error = graph
.validate()
.expect_err("a cyclic graph must be rejected");
assert!(matches!(error, GpuOptimError::InvalidState(_)));
assert!(FusionPlanner::new().plan(&graph).is_err());
}
#[test]
fn dangling_input_rejected() {
let mut graph = FusionGraph::new();
let _a = graph.add_op(OpKind::Relu, vec![99], vec![8], 4);
let error = graph
.validate()
.expect_err("a dangling input must be rejected");
assert!(matches!(error, GpuOptimError::InvalidState(_)));
}
#[test]
fn zero_dtype_rejected() {
let mut graph = FusionGraph::new();
let _a = graph.add_op(OpKind::Relu, vec![], vec![8], 0);
assert!(graph.validate().is_err());
}
#[test]
fn empty_graph_is_valid() {
let graph = FusionGraph::new();
assert!(graph.is_empty());
let plan = FusionPlanner::new()
.plan(&graph)
.expect("empty graph must plan");
assert_eq!(plan.num_groups(), 0);
assert_eq!(plan.bytes_unfused, 0);
assert_eq!(plan.bytes_fused, 0);
assert_eq!(plan.bytes_saved, 0);
assert!((plan.speedup_estimate - 1.0).abs() < 1e-9);
}
}