1use std::collections::{HashMap, HashSet};
29
30use crate::analysis::{dominance_analyse, topo_analyse};
31use crate::error::{GraphError, GraphResult};
32use crate::graph::ComputeGraph;
33use crate::node::{KernelConfig, NodeId, NodeKind};
34
35#[derive(Debug, Clone, PartialEq, Eq)]
44pub struct FusionGroup {
45 pub id: usize,
47 pub members: Vec<NodeId>,
49 pub config: KernelConfig,
51 pub tag: String,
53}
54
55impl FusionGroup {
56 #[must_use]
58 pub fn size(&self) -> usize {
59 self.members.len()
60 }
61
62 #[must_use]
64 pub fn is_trivial(&self) -> bool {
65 self.members.len() == 1
66 }
67}
68
69#[derive(Debug, Clone)]
75pub struct FusionPlan {
76 pub groups: Vec<FusionGroup>,
78 pub node_to_group: HashMap<NodeId, usize>,
80}
81
82impl FusionPlan {
83 pub fn fusion_count(&self) -> usize {
85 self.groups.iter().filter(|g| !g.is_trivial()).count()
86 }
87
88 pub fn nodes_saved(&self) -> usize {
92 self.groups
93 .iter()
94 .filter(|g| !g.is_trivial())
95 .map(|g| g.size() - 1)
96 .sum()
97 }
98
99 pub fn group_of(&self, node: NodeId) -> Option<&FusionGroup> {
101 self.node_to_group
102 .get(&node)
103 .and_then(|&idx| self.groups.get(idx))
104 }
105}
106
107fn configs_compatible(a: &KernelConfig, b: &KernelConfig) -> bool {
116 a.total_threads() == b.total_threads()
117}
118
119fn only_fusible_between(
122 graph: &ComputeGraph,
123 a: NodeId,
124 b: NodeId,
125 topo_pos: &HashMap<NodeId, usize>,
126) -> bool {
127 let pos_a = topo_pos[&a];
128 let pos_b = topo_pos[&b];
129 if pos_b <= pos_a + 1 {
130 return true; }
132 let mut visited = HashSet::new();
135 let mut stack = vec![a];
136 while let Some(cur) = stack.pop() {
137 if cur == b {
138 continue;
139 }
140 for &s in graph.successors(cur).unwrap_or(&[]) {
141 if visited.insert(s) {
142 if s == b {
143 continue;
144 }
145 let node = graph.node(s).ok();
146 let is_fusible = node.map(|n| n.kind.is_fusible()).unwrap_or(false);
147 let is_barrier = node
148 .map(|n| matches!(n.kind, NodeKind::Barrier))
149 .unwrap_or(false);
150 let spos = topo_pos.get(&s).copied().unwrap_or(usize::MAX);
151 if spos < pos_b && (is_fusible || is_barrier) {
152 stack.push(s);
153 } else if spos < pos_b && !is_fusible && !is_barrier {
154 return false; }
156 }
157 }
158 }
159 true
160}
161
162pub fn analyse(graph: &ComputeGraph) -> GraphResult<FusionPlan> {
174 if graph.is_empty() {
175 return Err(GraphError::EmptyGraph);
176 }
177
178 let topo = topo_analyse(graph)?;
179 let dt = dominance_analyse(graph)?;
180
181 let topo_pos: HashMap<NodeId, usize> = topo
182 .order
183 .iter()
184 .enumerate()
185 .map(|(p, &id)| (id, p))
186 .collect();
187
188 let mut assigned: HashMap<NodeId, usize> = HashMap::new();
190 let mut groups: Vec<FusionGroup> = Vec::new();
191
192 for &node_id in &topo.order {
194 if assigned.contains_key(&node_id) {
195 continue;
196 }
197
198 let node = graph.node(node_id)?;
199
200 let (is_fusible, base_config) = match &node.kind {
202 NodeKind::KernelLaunch {
203 fusible, config, ..
204 } => (*fusible, *config),
205 _ => {
206 let gid = groups.len();
208 groups.push(FusionGroup {
209 id: gid,
210 members: vec![node_id],
211 config: KernelConfig::linear(1, 1, 0),
212 tag: format!("non_kernel_{}", node.kind.tag()),
213 });
214 assigned.insert(node_id, gid);
215 continue;
216 }
217 };
218
219 if !is_fusible {
220 let gid = groups.len();
221 groups.push(FusionGroup {
222 id: gid,
223 members: vec![node_id],
224 config: base_config,
225 tag: format!("non_fusible_{}", node.display_name()),
226 });
227 assigned.insert(node_id, gid);
228 continue;
229 }
230
231 let gid = groups.len();
233 let mut members = vec![node_id];
234 assigned.insert(node_id, gid);
235
236 let mut frontier = graph.successors(node_id)?.to_vec();
238 while let Some(succ_id) = frontier.first().copied() {
239 frontier.remove(0);
240 if assigned.contains_key(&succ_id) {
241 continue;
242 }
243 let succ = graph.node(succ_id)?;
244 let (succ_fusible, succ_config) = match &succ.kind {
245 NodeKind::KernelLaunch {
246 fusible, config, ..
247 } => (*fusible, *config),
248 _ => continue,
249 };
250 if !succ_fusible {
251 continue;
252 }
253 if !configs_compatible(&base_config, &succ_config) {
256 continue;
257 }
258 let last_member = *members.last().ok_or_else(|| {
260 GraphError::Internal("fusion group members unexpectedly empty".into())
261 })?;
262 if !dt.dominates(last_member, succ_id) {
263 continue;
264 }
265 if !only_fusible_between(graph, last_member, succ_id, &topo_pos) {
267 continue;
268 }
269 members.push(succ_id);
271 assigned.insert(succ_id, gid);
272 for &next in graph.successors(succ_id)? {
274 if !assigned.contains_key(&next) {
275 frontier.push(next);
276 }
277 }
278 }
279
280 let tag = if members.len() > 1 {
281 format!(
282 "fused_{}..{}",
283 graph.node(members[0])?.display_name(),
284 graph
285 .node(*members.last().ok_or_else(|| {
286 GraphError::Internal("fusion group members unexpectedly empty".into())
287 })?)?
288 .display_name()
289 )
290 } else {
291 format!("solo_{}", graph.node(node_id)?.display_name())
292 };
293
294 groups.push(FusionGroup {
295 id: gid,
296 members,
297 config: base_config,
298 tag,
299 });
300 }
301
302 let node_to_group: HashMap<NodeId, usize> = assigned;
304
305 Ok(FusionPlan {
306 groups,
307 node_to_group,
308 })
309}
310
311#[cfg(test)]
316mod tests {
317 use super::*;
318 use crate::builder::GraphBuilder;
319 use crate::node::MemcpyDir;
320
321 fn fusible_kernel(b: &mut GraphBuilder, name: &str) -> NodeId {
322 b.add_kernel(name, 4, 256, 0).fusible(true).finish()
323 }
324
325 fn non_fusible_kernel(b: &mut GraphBuilder, name: &str) -> NodeId {
326 b.add_kernel(name, 4, 256, 0).fusible(false).finish()
327 }
328
329 #[test]
330 fn fusion_empty_graph() {
331 let g = ComputeGraph::new();
332 assert!(matches!(analyse(&g), Err(GraphError::EmptyGraph)));
333 }
334
335 #[test]
336 fn fusion_single_fusible_kernel_trivial_group() {
337 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
338 let k = fusible_kernel(&mut b, "add");
339 let g = b.build().expect("test graph builds successfully");
340 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
341 assert_eq!(plan.groups.len(), 1);
342 assert!(plan.groups[0].is_trivial());
343 assert_eq!(
344 plan.group_of(k)
345 .expect("kernel node has fusion group")
346 .members,
347 vec![k]
348 );
349 }
350
351 #[test]
352 fn fusion_chain_of_fusible_kernels_merged() {
353 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
354 let k0 = fusible_kernel(&mut b, "k0");
355 let k1 = fusible_kernel(&mut b, "k1");
356 let k2 = fusible_kernel(&mut b, "k2");
357 b.chain(&[k0, k1, k2]);
358 let g = b.build().expect("test graph builds successfully");
359 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
360 assert_eq!(plan.fusion_count(), 1);
362 let group = plan.group_of(k0).expect("k0 has fusion group");
363 assert_eq!(group.size(), 3);
364 assert!(group.members.contains(&k0));
365 assert!(group.members.contains(&k1));
366 assert!(group.members.contains(&k2));
367 }
368
369 #[test]
370 fn fusion_non_fusible_breaks_chain() {
371 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
373 let k0 = fusible_kernel(&mut b, "k0");
374 let k1 = non_fusible_kernel(&mut b, "k1");
375 let k2 = fusible_kernel(&mut b, "k2");
376 b.chain(&[k0, k1, k2]);
377 let g = b.build().expect("test graph builds successfully");
378 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
379 let g0 = plan.group_of(k0).expect("k0 has fusion group").id;
381 let g2 = plan.group_of(k2).expect("k2 has fusion group").id;
382 assert_ne!(g0, g2);
383 }
384
385 #[test]
386 fn fusion_memcpy_not_fused() {
387 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
388 let upload = b.add_memcpy("up", MemcpyDir::HostToDevice, 1024);
389 let k = fusible_kernel(&mut b, "k");
390 let download = b.add_memcpy("dn", MemcpyDir::DeviceToHost, 1024);
391 b.chain(&[upload, k, download]);
392 let g = b.build().expect("test graph builds successfully");
393 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
394 let gup = plan.group_of(upload).expect("upload node has fusion group");
396 let gdn = plan
397 .group_of(download)
398 .expect("download node has fusion group");
399 assert!(gup.is_trivial());
400 assert!(gdn.is_trivial());
401 }
402
403 #[test]
404 fn fusion_incompatible_configs_not_fused() {
405 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
408 let k0 = b.add_kernel("k0", 4, 256, 0).fusible(true).finish();
409 let k1 = b.add_kernel("k1", 8, 256, 0).fusible(true).finish();
410 b.dep(k0, k1);
411 let g = b.build().expect("test graph builds successfully");
412 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
413 let gk0 = plan.group_of(k0).expect("k0 has fusion group").id;
414 let gk1 = plan.group_of(k1).expect("k1 has fusion group").id;
415 assert_ne!(gk0, gk1);
416 }
417
418 #[test]
419 fn fusion_nodes_saved_count() {
420 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
421 let k0 = fusible_kernel(&mut b, "k0");
422 let k1 = fusible_kernel(&mut b, "k1");
423 let k2 = fusible_kernel(&mut b, "k2");
424 b.chain(&[k0, k1, k2]);
425 let g = b.build().expect("test graph builds successfully");
426 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
427 assert_eq!(plan.nodes_saved(), 2);
429 }
430
431 #[test]
432 fn fusion_plan_covers_all_nodes() {
433 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
434 let k0 = fusible_kernel(&mut b, "k0");
435 let k1 = non_fusible_kernel(&mut b, "k1");
436 let upload = b.add_memcpy("up", MemcpyDir::HostToDevice, 512);
437 b.chain(&[upload, k0, k1]);
438 let g = b.build().expect("test graph builds successfully");
439 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
440 let total: usize = plan.groups.iter().map(|g| g.size()).sum();
442 assert_eq!(total, 3);
443 assert!(plan.node_to_group.contains_key(&k0));
445 assert!(plan.node_to_group.contains_key(&k1));
446 assert!(plan.node_to_group.contains_key(&upload));
447 }
448
449 #[test]
450 fn fusion_parallel_branches_not_fused() {
451 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
453 let src = b.add_barrier("src");
454 let k0 = fusible_kernel(&mut b, "k0");
455 let k1 = fusible_kernel(&mut b, "k1");
456 b.fan_out(src, &[k0, k1]);
457 let g = b.build().expect("test graph builds successfully");
458 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
459 let gk0 = plan.group_of(k0).expect("k0 has fusion group").id;
461 let gk1 = plan.group_of(k1).expect("k1 has fusion group").id;
462 assert_ne!(gk0, gk1);
463 }
464
465 #[test]
466 fn fusion_group_tag_contains_names() {
467 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
468 let k0 = fusible_kernel(&mut b, "relu");
469 let k1 = fusible_kernel(&mut b, "scale");
470 b.dep(k0, k1);
471 let g = b.build().expect("test graph builds successfully");
472 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
473 let group = plan.group_of(k0).expect("k0 has fusion group");
474 assert!(!group.tag.is_empty());
475 }
476
477 #[test]
478 fn fusion_empty_fusible_graph_one_group() {
479 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
480 let k = fusible_kernel(&mut b, "solo");
481 let g = b.build().expect("test graph builds successfully");
482 let plan = analyse(&g).expect("kernel fusion analysis succeeds on valid graph");
483 assert_eq!(plan.fusion_count(), 0); assert_eq!(plan.nodes_saved(), 0);
485 assert_eq!(
486 plan.group_of(k)
487 .expect("kernel node has fusion group")
488 .size(),
489 1
490 );
491 }
492}