1use std::collections::HashMap;
22
23use crate::error::{GraphError, GraphResult};
24use crate::graph::ComputeGraph;
25use crate::node::{BufferId, NodeId};
26
27#[derive(Debug, Clone, PartialEq, Eq)]
36pub struct LiveInterval {
37 pub buf: BufferId,
39 pub def_pos: Option<usize>,
42 pub last_use_pos: Option<usize>,
45 pub size_bytes: usize,
47 pub external: bool,
49}
50
51impl LiveInterval {
52 #[must_use]
54 pub fn start(&self) -> usize {
55 self.def_pos.unwrap_or(0)
56 }
57
58 #[must_use]
60 pub fn end(&self) -> usize {
61 self.last_use_pos.unwrap_or(self.start())
62 }
63
64 #[must_use]
68 pub fn overlaps(&self, other: &Self) -> bool {
69 self.start() <= other.end() && other.start() <= self.end()
70 }
71
72 #[must_use]
74 pub fn is_dead(&self) -> bool {
75 self.last_use_pos.is_none() && self.def_pos.is_some()
76 }
77
78 #[must_use]
80 pub fn length(&self) -> usize {
81 self.end() - self.start()
82 }
83}
84
85#[derive(Debug, Clone)]
91pub struct LivenessAnalysis {
92 pub order: Vec<NodeId>,
94 intervals: HashMap<BufferId, LiveInterval>,
96}
97
98impl LivenessAnalysis {
99 #[must_use]
101 pub fn interval(&self, buf: BufferId) -> Option<&LiveInterval> {
102 self.intervals.get(&buf)
103 }
104
105 pub fn all_intervals(&self) -> impl Iterator<Item = &LiveInterval> {
107 self.intervals.values()
108 }
109
110 pub fn sorted_by_start(&self) -> Vec<&LiveInterval> {
112 let mut ivs: Vec<&LiveInterval> = self.intervals.values().collect();
113 ivs.sort_by_key(|i| (i.start(), i.buf.0));
114 ivs
115 }
116
117 pub fn interference_pairs(&self) -> Vec<(BufferId, BufferId)> {
122 let ivs: Vec<&LiveInterval> = self.intervals.values().collect();
123 let mut pairs = Vec::new();
124 for i in 0..ivs.len() {
125 for j in (i + 1)..ivs.len() {
126 if ivs[i].overlaps(ivs[j]) {
127 let a = ivs[i].buf.min(ivs[j].buf);
128 let b = ivs[i].buf.max(ivs[j].buf);
129 pairs.push((a, b));
130 }
131 }
132 }
133 pairs.sort();
134 pairs.dedup();
135 pairs
136 }
137
138 pub fn dead_buffers(&self) -> Vec<BufferId> {
140 self.intervals
141 .values()
142 .filter(|i| i.is_dead())
143 .map(|i| i.buf)
144 .collect()
145 }
146
147 pub fn max_live_count(&self) -> usize {
151 if self.order.is_empty() {
152 return 0;
153 }
154 let n_steps = self.order.len();
155 let mut count_at = vec![0usize; n_steps];
156 for iv in self.intervals.values() {
157 let range = iv.start()..=iv.end().min(n_steps - 1);
158 for cnt in count_at[range].iter_mut() {
159 *cnt += 1;
160 }
161 }
162 *count_at.iter().max().unwrap_or(&0)
163 }
164
165 pub fn max_live_bytes(&self) -> usize {
167 if self.order.is_empty() {
168 return 0;
169 }
170 let n_steps = self.order.len();
171 let mut bytes_at = vec![0usize; n_steps];
172 for iv in self.intervals.values() {
173 if iv.external {
174 continue; }
176 let range = iv.start()..=iv.end().min(n_steps - 1);
177 let sz = iv.size_bytes;
178 for b in bytes_at[range].iter_mut() {
179 *b = b.saturating_add(sz);
180 }
181 }
182 *bytes_at.iter().max().unwrap_or(&0)
183 }
184}
185
186pub fn analyse(graph: &ComputeGraph) -> GraphResult<LivenessAnalysis> {
196 if graph.is_empty() {
197 return Err(GraphError::EmptyGraph);
198 }
199
200 let order = graph.topological_order()?;
201
202 let pos_of: HashMap<NodeId, usize> = order.iter().enumerate().map(|(p, &id)| (id, p)).collect();
204
205 let mut intervals: HashMap<BufferId, LiveInterval> = HashMap::new();
206
207 for buf in graph.buffers() {
209 intervals.insert(
210 buf.id,
211 LiveInterval {
212 buf: buf.id,
213 def_pos: None,
214 last_use_pos: None,
215 size_bytes: buf.size_bytes,
216 external: buf.external,
217 },
218 );
219 }
220
221 for &node_id in &order {
223 let node = graph.node(node_id)?;
224 let p = pos_of[&node_id];
225
226 for &buf in &node.outputs {
228 let iv = intervals.entry(buf).or_insert_with(|| LiveInterval {
229 buf,
230 def_pos: None,
231 last_use_pos: None,
232 size_bytes: 0,
233 external: false,
234 });
235 if iv.def_pos.is_none() {
237 iv.def_pos = Some(p);
238 }
239 }
240
241 for &buf in &node.inputs {
243 let iv = intervals.entry(buf).or_insert_with(|| LiveInterval {
244 buf,
245 def_pos: None,
246 last_use_pos: None,
247 size_bytes: 0,
248 external: false,
249 });
250 match iv.last_use_pos {
251 None => iv.last_use_pos = Some(p),
252 Some(prev) if p > prev => iv.last_use_pos = Some(p),
253 _ => {}
254 }
255 }
256 }
257
258 Ok(LivenessAnalysis { order, intervals })
259}
260
261#[cfg(test)]
266mod tests {
267 use super::*;
268 use crate::builder::GraphBuilder;
269
270 fn build_linear_graph() -> (ComputeGraph, BufferId, NodeId, NodeId) {
271 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
273 let buf = b.alloc_buffer("shared", 1024);
274 let writer = b.add_barrier("writer");
275 let reader = b.add_barrier("reader");
276 b.set_outputs(writer, [buf]);
277 b.set_inputs(reader, [buf]);
278 b.dep(writer, reader);
279 let g = b.build().expect("test graph builds successfully");
280 (g, buf, writer, reader)
281 }
282
283 #[test]
284 fn liveness_empty_graph() {
285 let g = ComputeGraph::new();
286 assert!(matches!(analyse(&g), Err(GraphError::EmptyGraph)));
287 }
288
289 #[test]
290 fn liveness_buffer_def_and_use() {
291 let (g, buf, writer, reader) = build_linear_graph();
292 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
293 let iv = la
294 .interval(buf)
295 .expect("buffer registered in liveness analysis");
296 let order = &la.order;
297 let wpos = order
298 .iter()
299 .position(|&x| x == writer)
300 .expect("writer node present in topological order");
301 let rpos = order
302 .iter()
303 .position(|&x| x == reader)
304 .expect("reader node present in topological order");
305 assert_eq!(iv.def_pos, Some(wpos));
306 assert_eq!(iv.last_use_pos, Some(rpos));
307 assert!(!iv.is_dead());
308 }
309
310 #[test]
311 fn liveness_dead_buffer() {
312 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
313 let buf = b.alloc_buffer("dead", 512);
314 let writer = b.add_barrier("w");
315 b.set_outputs(writer, [buf]);
316 let g = b.build().expect("test graph builds successfully");
317 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
318 let iv = la
319 .interval(buf)
320 .expect("buffer registered in liveness analysis");
321 assert!(iv.is_dead());
322 }
323
324 #[test]
325 fn liveness_external_buffer_not_counted_in_bytes() {
326 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
327 let ext = b.alloc_external_buffer("ext", 65536);
328 let reader = b.add_barrier("r");
329 b.set_inputs(reader, [ext]);
330 let g = b.build().expect("test graph builds successfully");
331 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
332 assert_eq!(la.max_live_bytes(), 0);
334 }
335
336 #[test]
337 fn liveness_overlap_detection() {
338 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
340 let buf0 = b.alloc_buffer("b0", 1024);
341 let buf1 = b.alloc_buffer("b1", 2048);
342 let a = b.add_barrier("a");
343 let bnode = b.add_barrier("b");
344 let c = b.add_barrier("c");
345 b.set_outputs(a, [buf0]);
346 b.set_outputs(bnode, [buf1]);
347 b.set_inputs(c, [buf0, buf1]);
348 b.dep(a, c).dep(bnode, c);
349 let g = b.build().expect("test graph builds successfully");
350 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
351 let pairs = la.interference_pairs();
352 assert!(pairs.contains(&(BufferId(0), BufferId(1))));
354 }
355
356 #[test]
357 fn liveness_non_overlapping_no_interference() {
358 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
362 let buf0 = b.alloc_buffer("b0", 512);
363 let buf1 = b.alloc_buffer("b1", 512);
364 let n0 = b.add_barrier("n0");
369 let n1 = b.add_barrier("n1");
370 let n2 = b.add_barrier("n2");
371 let n3 = b.add_barrier("n3");
372 b.set_outputs(n0, [buf0]);
373 b.set_inputs(n1, [buf0]);
374 b.set_outputs(n2, [buf1]);
375 b.set_inputs(n3, [buf1]);
376 b.chain(&[n0, n1, n2, n3]);
377 let g = b.build().expect("test graph builds successfully");
378 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
379 let iv0 = la
381 .interval(buf0)
382 .expect("buf0 registered in liveness analysis");
383 let iv1 = la
384 .interval(buf1)
385 .expect("buf1 registered in liveness analysis");
386 assert!(!iv0.overlaps(iv1));
387 assert!(la.interference_pairs().is_empty());
388 }
389
390 #[test]
391 fn liveness_max_live_count() {
392 let (g, _buf, _w, _r) = build_linear_graph();
393 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
394 assert_eq!(la.max_live_count(), 1);
396 }
397
398 #[test]
399 fn liveness_max_live_bytes() {
400 let (g, _buf, _w, _r) = build_linear_graph();
401 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
402 assert_eq!(la.max_live_bytes(), 1024);
404 }
405
406 #[test]
407 fn liveness_sorted_by_start() {
408 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
409 let buf0 = b.alloc_buffer("b0", 100);
410 let buf1 = b.alloc_buffer("b1", 200);
411 let n0 = b.add_barrier("n0");
412 let n1 = b.add_barrier("n1");
413 b.set_outputs(n0, [buf0]);
414 b.set_outputs(n1, [buf1]);
415 b.dep(n0, n1);
416 let g = b.build().expect("test graph builds successfully");
417 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
418 let sorted = la.sorted_by_start();
419 if sorted.len() == 2 {
421 assert!(sorted[0].start() <= sorted[1].start());
422 }
423 }
424
425 #[test]
426 fn liveness_interval_length() {
427 let (g, buf, _w, _r) = build_linear_graph();
428 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
429 let iv = la
430 .interval(buf)
431 .expect("buffer registered in liveness analysis");
432 assert_eq!(iv.length(), 1);
434 }
435
436 #[test]
437 fn liveness_dead_buffers_list() {
438 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
439 let dead = b.alloc_buffer("dead", 1);
440 let _w = {
441 let w = b.add_barrier("w");
442 b.set_outputs(w, [dead]);
443 w
444 };
445 let g = b.build().expect("test graph builds successfully");
446 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
447 let dead_list = la.dead_buffers();
448 assert!(dead_list.contains(&dead));
449 }
450
451 #[test]
452 fn liveness_all_intervals_count() {
453 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
454 b.alloc_buffer("a", 1);
455 b.alloc_buffer("b", 2);
456 b.alloc_buffer("c", 4);
457 b.add_barrier("n");
458 let g = b.build().expect("test graph builds successfully");
459 let la = analyse(&g).expect("liveness analysis succeeds on valid graph");
460 assert_eq!(la.all_intervals().count(), 3);
462 }
463}