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().unwrap();
340 let plan = analyse(&g).unwrap();
341 assert_eq!(plan.groups.len(), 1);
342 assert!(plan.groups[0].is_trivial());
343 assert_eq!(plan.group_of(k).unwrap().members, vec![k]);
344 }
345
346 #[test]
347 fn fusion_chain_of_fusible_kernels_merged() {
348 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
349 let k0 = fusible_kernel(&mut b, "k0");
350 let k1 = fusible_kernel(&mut b, "k1");
351 let k2 = fusible_kernel(&mut b, "k2");
352 b.chain(&[k0, k1, k2]);
353 let g = b.build().unwrap();
354 let plan = analyse(&g).unwrap();
355 assert_eq!(plan.fusion_count(), 1);
357 let group = plan.group_of(k0).unwrap();
358 assert_eq!(group.size(), 3);
359 assert!(group.members.contains(&k0));
360 assert!(group.members.contains(&k1));
361 assert!(group.members.contains(&k2));
362 }
363
364 #[test]
365 fn fusion_non_fusible_breaks_chain() {
366 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
368 let k0 = fusible_kernel(&mut b, "k0");
369 let k1 = non_fusible_kernel(&mut b, "k1");
370 let k2 = fusible_kernel(&mut b, "k2");
371 b.chain(&[k0, k1, k2]);
372 let g = b.build().unwrap();
373 let plan = analyse(&g).unwrap();
374 let g0 = plan.group_of(k0).unwrap().id;
376 let g2 = plan.group_of(k2).unwrap().id;
377 assert_ne!(g0, g2);
378 }
379
380 #[test]
381 fn fusion_memcpy_not_fused() {
382 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
383 let upload = b.add_memcpy("up", MemcpyDir::HostToDevice, 1024);
384 let k = fusible_kernel(&mut b, "k");
385 let download = b.add_memcpy("dn", MemcpyDir::DeviceToHost, 1024);
386 b.chain(&[upload, k, download]);
387 let g = b.build().unwrap();
388 let plan = analyse(&g).unwrap();
389 let gup = plan.group_of(upload).unwrap();
391 let gdn = plan.group_of(download).unwrap();
392 assert!(gup.is_trivial());
393 assert!(gdn.is_trivial());
394 }
395
396 #[test]
397 fn fusion_incompatible_configs_not_fused() {
398 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
401 let k0 = b.add_kernel("k0", 4, 256, 0).fusible(true).finish();
402 let k1 = b.add_kernel("k1", 8, 256, 0).fusible(true).finish();
403 b.dep(k0, k1);
404 let g = b.build().unwrap();
405 let plan = analyse(&g).unwrap();
406 let gk0 = plan.group_of(k0).unwrap().id;
407 let gk1 = plan.group_of(k1).unwrap().id;
408 assert_ne!(gk0, gk1);
409 }
410
411 #[test]
412 fn fusion_nodes_saved_count() {
413 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
414 let k0 = fusible_kernel(&mut b, "k0");
415 let k1 = fusible_kernel(&mut b, "k1");
416 let k2 = fusible_kernel(&mut b, "k2");
417 b.chain(&[k0, k1, k2]);
418 let g = b.build().unwrap();
419 let plan = analyse(&g).unwrap();
420 assert_eq!(plan.nodes_saved(), 2);
422 }
423
424 #[test]
425 fn fusion_plan_covers_all_nodes() {
426 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
427 let k0 = fusible_kernel(&mut b, "k0");
428 let k1 = non_fusible_kernel(&mut b, "k1");
429 let upload = b.add_memcpy("up", MemcpyDir::HostToDevice, 512);
430 b.chain(&[upload, k0, k1]);
431 let g = b.build().unwrap();
432 let plan = analyse(&g).unwrap();
433 let total: usize = plan.groups.iter().map(|g| g.size()).sum();
435 assert_eq!(total, 3);
436 assert!(plan.node_to_group.contains_key(&k0));
438 assert!(plan.node_to_group.contains_key(&k1));
439 assert!(plan.node_to_group.contains_key(&upload));
440 }
441
442 #[test]
443 fn fusion_parallel_branches_not_fused() {
444 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
446 let src = b.add_barrier("src");
447 let k0 = fusible_kernel(&mut b, "k0");
448 let k1 = fusible_kernel(&mut b, "k1");
449 b.fan_out(src, &[k0, k1]);
450 let g = b.build().unwrap();
451 let plan = analyse(&g).unwrap();
452 let gk0 = plan.group_of(k0).unwrap().id;
454 let gk1 = plan.group_of(k1).unwrap().id;
455 assert_ne!(gk0, gk1);
456 }
457
458 #[test]
459 fn fusion_group_tag_contains_names() {
460 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
461 let k0 = fusible_kernel(&mut b, "relu");
462 let k1 = fusible_kernel(&mut b, "scale");
463 b.dep(k0, k1);
464 let g = b.build().unwrap();
465 let plan = analyse(&g).unwrap();
466 let group = plan.group_of(k0).unwrap();
467 assert!(!group.tag.is_empty());
468 }
469
470 #[test]
471 fn fusion_empty_fusible_graph_one_group() {
472 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
473 let k = fusible_kernel(&mut b, "solo");
474 let g = b.build().unwrap();
475 let plan = analyse(&g).unwrap();
476 assert_eq!(plan.fusion_count(), 0); assert_eq!(plan.nodes_saved(), 0);
478 assert_eq!(plan.group_of(k).unwrap().size(), 1);
479 }
480}