Skip to main content

candle_graph/
timing.rs

1//! Separate host and device timing planes with overlap-safe device aggregation.
2
3use std::collections::BTreeMap;
4
5use serde::{Deserialize, Serialize};
6
7use crate::capability::CoverageLevel;
8use crate::trace::TraceDocument;
9
10#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
11pub struct DeviceSpanTiming {
12    pub span_id: String,
13    pub device: String,
14    pub clock_id: String,
15    pub backends: Vec<String>,
16    pub streams: Vec<String>,
17    pub interval_count: usize,
18    /// Union of all intervals for this span on this device clock.
19    pub busy_ns: u64,
20}
21
22#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
23pub struct DeviceClockTiming {
24    pub device: String,
25    pub clock_id: String,
26    pub interval_count: usize,
27    /// Union across streams; overlapping device work is counted once.
28    pub busy_ns: u64,
29}
30
31#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
32pub struct TimingProfile {
33    pub device_coverage: CoverageLevel,
34    pub device_spans: Vec<DeviceSpanTiming>,
35    pub device_clocks: Vec<DeviceClockTiming>,
36}
37
38pub fn analyze_timing(doc: &TraceDocument) -> TimingProfile {
39    let mut by_span: BTreeMap<(String, String, String), Vec<_>> = BTreeMap::new();
40    let mut by_clock: BTreeMap<(String, String), Vec<_>> = BTreeMap::new();
41    for interval in &doc.device_intervals {
42        by_span
43            .entry((
44                interval.span_id.clone(),
45                interval.device.clone(),
46                interval.clock_id.clone(),
47            ))
48            .or_default()
49            .push(interval);
50        by_clock
51            .entry((interval.device.clone(), interval.clock_id.clone()))
52            .or_default()
53            .push(interval);
54    }
55
56    let device_spans = by_span
57        .into_iter()
58        .map(|((span_id, device, clock_id), intervals)| {
59            let mut backends = intervals
60                .iter()
61                .map(|interval| interval.backend.clone())
62                .collect::<Vec<_>>();
63            backends.sort();
64            backends.dedup();
65            let mut streams = intervals
66                .iter()
67                .map(|interval| interval.stream_id.clone())
68                .collect::<Vec<_>>();
69            streams.sort();
70            streams.dedup();
71            DeviceSpanTiming {
72                span_id,
73                device,
74                clock_id,
75                backends,
76                streams,
77                interval_count: intervals.len(),
78                busy_ns: interval_union_ns(intervals.iter().map(|item| {
79                    (
80                        item.start_ns,
81                        item.start_ns.saturating_add(item.duration_ns),
82                    )
83                })),
84            }
85        })
86        .collect();
87
88    let device_clocks = by_clock
89        .into_iter()
90        .map(|((device, clock_id), intervals)| DeviceClockTiming {
91            device,
92            clock_id,
93            interval_count: intervals.len(),
94            busy_ns: interval_union_ns(intervals.iter().map(|item| {
95                (
96                    item.start_ns,
97                    item.start_ns.saturating_add(item.duration_ns),
98                )
99            })),
100        })
101        .collect();
102
103    TimingProfile {
104        device_coverage: doc
105            .run
106            .capture_contract
107            .device_timing
108            .with_observations(doc.device_intervals.len()),
109        device_spans,
110        device_clocks,
111    }
112}
113
114fn interval_union_ns(intervals: impl IntoIterator<Item = (u64, u64)>) -> u64 {
115    let mut intervals = intervals.into_iter().collect::<Vec<_>>();
116    intervals.sort_unstable();
117    let mut total = 0u64;
118    let mut current: Option<(u64, u64)> = None;
119    for (start, end) in intervals {
120        match current {
121            None => current = Some((start, end)),
122            Some((left, right)) if start <= right => current = Some((left, right.max(end))),
123            Some((left, right)) => {
124                total = total.saturating_add(right.saturating_sub(left));
125                current = Some((start, end));
126            }
127        }
128    }
129    if let Some((start, end)) = current {
130        total = total.saturating_add(end.saturating_sub(start));
131    }
132    total
133}
134
135#[cfg(test)]
136mod tests {
137    use super::interval_union_ns;
138
139    #[test]
140    fn overlapping_intervals_are_counted_once() {
141        assert_eq!(interval_union_ns([(0, 10), (5, 20), (30, 35)]), 25);
142    }
143}