Skip to main content

oxicuda_graph/analysis/
liveness.rs

1//! Buffer liveness analysis.
2//!
3//! For each buffer in a `ComputeGraph`, this pass computes the **live
4//! interval** `[def_pos, last_use_pos]` in terms of positions in the
5//! topological order. The memory planner uses these intervals to determine
6//! which buffers can share the same physical memory allocation.
7//!
8//! # Algorithm
9//!
10//! 1. Run a topological sort to get a linear order of nodes.
11//! 2. For each node at position `p`:
12//!    - For every output buffer `b` it writes: record `def_pos[b] = p`.
13//!    - For every input buffer `b` it reads: record `last_use_pos[b] = max(last_use_pos[b], p)`.
14//! 3. The live interval for buffer `b` is `[def_pos[b], last_use_pos[b]]`.
15//!    Buffers that are only written (no reads) have `last_use_pos = def_pos`.
16//!    Buffers that are only read (external inputs) have `def_pos = 0`.
17//!
18//! Two buffers `a` and `b` **interfere** if their live intervals overlap, i.e.
19//! `a.def <= b.last_use && b.def <= a.last_use`.
20
21use std::collections::HashMap;
22
23use crate::error::{GraphError, GraphResult};
24use crate::graph::ComputeGraph;
25use crate::node::{BufferId, NodeId};
26
27// ---------------------------------------------------------------------------
28// LiveInterval
29// ---------------------------------------------------------------------------
30
31/// The live interval of a buffer in topological-order position space.
32///
33/// The buffer is "live" from the step it is first written (`def_pos`) to the
34/// step at which it is last read (`last_use_pos`), inclusive.
35#[derive(Debug, Clone, PartialEq, Eq)]
36pub struct LiveInterval {
37    /// Buffer this interval belongs to.
38    pub buf: BufferId,
39    /// Topological position at which the buffer is first defined (written).
40    /// `None` means the buffer is an external input (alive from the start).
41    pub def_pos: Option<usize>,
42    /// Topological position at which the buffer is last used (read).
43    /// `None` means the buffer is never read (write-only / dead code).
44    pub last_use_pos: Option<usize>,
45    /// Buffer size in bytes (copied from `BufferDescriptor`).
46    pub size_bytes: usize,
47    /// Whether this buffer is externally managed.
48    pub external: bool,
49}
50
51impl LiveInterval {
52    /// Returns the effective start position (0 if never written).
53    #[must_use]
54    pub fn start(&self) -> usize {
55        self.def_pos.unwrap_or(0)
56    }
57
58    /// Returns the effective end position (start if never read).
59    #[must_use]
60    pub fn end(&self) -> usize {
61        self.last_use_pos.unwrap_or(self.start())
62    }
63
64    /// Returns `true` if this interval overlaps with `other`.
65    ///
66    /// Two intervals overlap when neither ends strictly before the other begins.
67    #[must_use]
68    pub fn overlaps(&self, other: &Self) -> bool {
69        self.start() <= other.end() && other.start() <= self.end()
70    }
71
72    /// Returns `true` if the buffer is dead (never read after being written).
73    #[must_use]
74    pub fn is_dead(&self) -> bool {
75        self.last_use_pos.is_none() && self.def_pos.is_some()
76    }
77
78    /// Returns the length of the live interval in steps.
79    #[must_use]
80    pub fn length(&self) -> usize {
81        self.end() - self.start()
82    }
83}
84
85// ---------------------------------------------------------------------------
86// LivenessAnalysis
87// ---------------------------------------------------------------------------
88
89/// Result of running liveness analysis on a `ComputeGraph`.
90#[derive(Debug, Clone)]
91pub struct LivenessAnalysis {
92    /// The topological order used for position assignment.
93    pub order: Vec<NodeId>,
94    /// Live intervals keyed by buffer ID.
95    intervals: HashMap<BufferId, LiveInterval>,
96}
97
98impl LivenessAnalysis {
99    /// Returns the live interval for a buffer, if it was referenced.
100    #[must_use]
101    pub fn interval(&self, buf: BufferId) -> Option<&LiveInterval> {
102        self.intervals.get(&buf)
103    }
104
105    /// Returns all live intervals.
106    pub fn all_intervals(&self) -> impl Iterator<Item = &LiveInterval> {
107        self.intervals.values()
108    }
109
110    /// Returns all intervals sorted by start position (ascending).
111    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    /// Returns pairs of buffers whose live intervals overlap (interference set).
118    ///
119    /// The result is deduplicated: each pair `(a, b)` appears at most once,
120    /// with `a.0 < b.0`.
121    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    /// Returns dead buffers (written but never read).
139    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    /// Returns the maximum number of buffers simultaneously live at any step.
148    ///
149    /// This is a lower bound on the number of live allocations required.
150    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    /// Returns the maximum total bytes simultaneously live at any step.
166    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; // externally managed, not counted in device memory
175            }
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
186// ---------------------------------------------------------------------------
187// analyse — entry point
188// ---------------------------------------------------------------------------
189
190/// Computes live intervals for all buffers referenced in `graph`.
191///
192/// # Errors
193///
194/// Returns [`GraphError::EmptyGraph`] if the graph has no nodes.
195pub 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    // Map from position in topological order.
203    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    // Initialise entries for all registered buffers.
208    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    // Scan nodes in topological order.
222    for &node_id in &order {
223        let node = graph.node(node_id)?;
224        let p = pos_of[&node_id];
225
226        // Outputs: this node defines (writes) the buffer.
227        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            // Only record the first definition.
236            if iv.def_pos.is_none() {
237                iv.def_pos = Some(p);
238            }
239        }
240
241        // Inputs: this node uses (reads) the buffer; update last-use.
242        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// ---------------------------------------------------------------------------
262// Tests
263// ---------------------------------------------------------------------------
264
265#[cfg(test)]
266mod tests {
267    use super::*;
268    use crate::builder::GraphBuilder;
269
270    fn build_linear_graph() -> (ComputeGraph, BufferId, NodeId, NodeId) {
271        // writer → reader, sharing buffer `buf`.
272        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        // External buffer should not count toward live bytes.
333        assert_eq!(la.max_live_bytes(), 0);
334    }
335
336    #[test]
337    fn liveness_overlap_detection() {
338        // a writes buf0; b writes buf1; c reads buf0 and buf1.
339        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        // buf0 and buf1 both live until c, so they interfere.
353        assert!(pairs.contains(&(BufferId(0), BufferId(1))));
354    }
355
356    #[test]
357    fn liveness_non_overlapping_no_interference() {
358        // a writes buf0, b reads buf0, b writes buf1, c reads buf1 — chain.
359        // buf0 is live [0,1], buf1 is live [1,2] — they share position 1 so they DO overlap.
360        // Let's test complete disjoint: buf0 live [0,0], buf1 live [2,2].
361        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        // n0 writes buf0 (live start=0)
365        // n1 reads buf0 (live end=1) and writes nothing
366        // n2 writes buf1 (live start=2)
367        // n3 reads buf1 (live end=3)
368        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        // buf0: [0,1], buf1: [2,3] — no overlap.
380        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        // Only one buffer, so max live = 1.
395        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        // One internal buffer of 1024 bytes.
403        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        // buf0 defined at step 0, buf1 at step 1 → buf0 first.
420        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        // writer at 0, reader at 1 → length = 1.
433        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        // 3 buffers registered.
461        assert_eq!(la.all_intervals().count(), 3);
462    }
463}