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().unwrap();
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).unwrap();
293 let iv = la.interval(buf).unwrap();
294 let order = &la.order;
295 let wpos = order.iter().position(|&x| x == writer).unwrap();
296 let rpos = order.iter().position(|&x| x == reader).unwrap();
297 assert_eq!(iv.def_pos, Some(wpos));
298 assert_eq!(iv.last_use_pos, Some(rpos));
299 assert!(!iv.is_dead());
300 }
301
302 #[test]
303 fn liveness_dead_buffer() {
304 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
305 let buf = b.alloc_buffer("dead", 512);
306 let writer = b.add_barrier("w");
307 b.set_outputs(writer, [buf]);
308 let g = b.build().unwrap();
309 let la = analyse(&g).unwrap();
310 let iv = la.interval(buf).unwrap();
311 assert!(iv.is_dead());
312 }
313
314 #[test]
315 fn liveness_external_buffer_not_counted_in_bytes() {
316 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
317 let ext = b.alloc_external_buffer("ext", 65536);
318 let reader = b.add_barrier("r");
319 b.set_inputs(reader, [ext]);
320 let g = b.build().unwrap();
321 let la = analyse(&g).unwrap();
322 assert_eq!(la.max_live_bytes(), 0);
324 }
325
326 #[test]
327 fn liveness_overlap_detection() {
328 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
330 let buf0 = b.alloc_buffer("b0", 1024);
331 let buf1 = b.alloc_buffer("b1", 2048);
332 let a = b.add_barrier("a");
333 let bnode = b.add_barrier("b");
334 let c = b.add_barrier("c");
335 b.set_outputs(a, [buf0]);
336 b.set_outputs(bnode, [buf1]);
337 b.set_inputs(c, [buf0, buf1]);
338 b.dep(a, c).dep(bnode, c);
339 let g = b.build().unwrap();
340 let la = analyse(&g).unwrap();
341 let pairs = la.interference_pairs();
342 assert!(pairs.contains(&(BufferId(0), BufferId(1))));
344 }
345
346 #[test]
347 fn liveness_non_overlapping_no_interference() {
348 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
352 let buf0 = b.alloc_buffer("b0", 512);
353 let buf1 = b.alloc_buffer("b1", 512);
354 let n0 = b.add_barrier("n0");
359 let n1 = b.add_barrier("n1");
360 let n2 = b.add_barrier("n2");
361 let n3 = b.add_barrier("n3");
362 b.set_outputs(n0, [buf0]);
363 b.set_inputs(n1, [buf0]);
364 b.set_outputs(n2, [buf1]);
365 b.set_inputs(n3, [buf1]);
366 b.chain(&[n0, n1, n2, n3]);
367 let g = b.build().unwrap();
368 let la = analyse(&g).unwrap();
369 let iv0 = la.interval(buf0).unwrap();
371 let iv1 = la.interval(buf1).unwrap();
372 assert!(!iv0.overlaps(iv1));
373 assert!(la.interference_pairs().is_empty());
374 }
375
376 #[test]
377 fn liveness_max_live_count() {
378 let (g, _buf, _w, _r) = build_linear_graph();
379 let la = analyse(&g).unwrap();
380 assert_eq!(la.max_live_count(), 1);
382 }
383
384 #[test]
385 fn liveness_max_live_bytes() {
386 let (g, _buf, _w, _r) = build_linear_graph();
387 let la = analyse(&g).unwrap();
388 assert_eq!(la.max_live_bytes(), 1024);
390 }
391
392 #[test]
393 fn liveness_sorted_by_start() {
394 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
395 let buf0 = b.alloc_buffer("b0", 100);
396 let buf1 = b.alloc_buffer("b1", 200);
397 let n0 = b.add_barrier("n0");
398 let n1 = b.add_barrier("n1");
399 b.set_outputs(n0, [buf0]);
400 b.set_outputs(n1, [buf1]);
401 b.dep(n0, n1);
402 let g = b.build().unwrap();
403 let la = analyse(&g).unwrap();
404 let sorted = la.sorted_by_start();
405 if sorted.len() == 2 {
407 assert!(sorted[0].start() <= sorted[1].start());
408 }
409 }
410
411 #[test]
412 fn liveness_interval_length() {
413 let (g, buf, _w, _r) = build_linear_graph();
414 let la = analyse(&g).unwrap();
415 let iv = la.interval(buf).unwrap();
416 assert_eq!(iv.length(), 1);
418 }
419
420 #[test]
421 fn liveness_dead_buffers_list() {
422 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
423 let dead = b.alloc_buffer("dead", 1);
424 let _w = {
425 let w = b.add_barrier("w");
426 b.set_outputs(w, [dead]);
427 w
428 };
429 let g = b.build().unwrap();
430 let la = analyse(&g).unwrap();
431 let dead_list = la.dead_buffers();
432 assert!(dead_list.contains(&dead));
433 }
434
435 #[test]
436 fn liveness_all_intervals_count() {
437 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
438 b.alloc_buffer("a", 1);
439 b.alloc_buffer("b", 2);
440 b.alloc_buffer("c", 4);
441 b.add_barrier("n");
442 let g = b.build().unwrap();
443 let la = analyse(&g).unwrap();
444 assert_eq!(la.all_intervals().count(), 3);
446 }
447}