Skip to main content

cranpose_render_wgpu/
pass_timing.rs

1use std::{
2    cell::{Cell, RefCell},
3    sync::{
4        Arc,
5        atomic::{AtomicU8, Ordering},
6    },
7};
8
9use crate::debug_toggles::DebugToggle;
10
11const QUERY_CAPACITY: u32 = 512;
12
13const READBACK_SLOTS: usize = 4;
14
15const SLOT_FREE: u8 = 0;
16const SLOT_PENDING: u8 = 1;
17const SLOT_MAPPED: u8 = 2;
18const SLOT_FAILED: u8 = 3;
19
20const PRINT_CADENCE_FRAMES: u64 = 60;
21
22static PASS_TIMING: DebugToggle = DebugToggle::new("CRANPOSE_GPU_PASS_TIMING");
23
24pub(crate) fn pass_timing_requested() -> bool {
25    PASS_TIMING.flag()
26}
27
28/// One label's aggregate inside the current print window.
29#[derive(Clone, Debug, PartialEq)]
30pub struct GpuPassTimingEntry {
31    pub label: String,
32    pub total_ms: f64,
33    pub passes: u64,
34}
35
36/// GPU time by pass label, aggregated since the last `[GPU-PASS]` print.
37#[derive(Clone, Debug, Default, PartialEq)]
38pub struct GpuPassTimingReport {
39    /// Frames whose timestamps have been read back into this window.
40    pub frames: u32,
41    /// Total GPU busy span — first pass begin to last pass end, summed over
42    /// the window's frames. Per-label times are stage-boundary occupancy
43    /// windows that overlap on pipelined GPUs and can sum past the frame;
44    /// this span is the frame-level wall number they cannot give.
45    pub span_ms: f64,
46    /// Entries sorted by descending GPU time; zero-time labels are omitted.
47    pub entries: Vec<GpuPassTimingEntry>,
48}
49
50#[derive(Clone, Copy, Default)]
51struct LabelTotal {
52    nanoseconds: u64,
53    passes: u64,
54}
55
56struct ReadbackSlot {
57    buffer: wgpu::Buffer,
58    state: Arc<AtomicU8>,
59    passes: RefCell<Vec<(u16, u32)>>,
60}
61
62pub(crate) struct PassTimer {
63    query_set: wgpu::QuerySet,
64    resolve_buffer: wgpu::Buffer,
65    period_ns: f32,
66    cursor: Cell<u32>,
67    frame_passes: RefCell<Vec<(u16, u32)>>,
68    labels: RefCell<Vec<String>>,
69    totals: RefCell<Vec<LabelTotal>>,
70    slots: Vec<ReadbackSlot>,
71    frame_index: Cell<u64>,
72    frames_harvested: Cell<u32>,
73    span_nanoseconds: Cell<u64>,
74    dropped_passes: Cell<u64>,
75    dropped_frames: Cell<u64>,
76}
77
78impl PassTimer {
79    pub(crate) fn for_device(device: &wgpu::Device, queue: &wgpu::Queue) -> Option<Self> {
80        if !device.features().contains(wgpu::Features::TIMESTAMP_QUERY) {
81            eprintln!(
82                "[GPU-PASS] CRANPOSE_GPU_PASS_TIMING is set but the adapter lacks TIMESTAMP_QUERY; passes will not be timed"
83            );
84            return None;
85        }
86        let query_set = device.create_query_set(&wgpu::QuerySetDescriptor {
87            label: Some("Pass Timing Query Set"),
88            ty: wgpu::QueryType::Timestamp,
89            count: QUERY_CAPACITY,
90        });
91        let buffer_size = u64::from(QUERY_CAPACITY) * u64::from(wgpu::QUERY_SIZE);
92        let resolve_buffer = device.create_buffer(&wgpu::BufferDescriptor {
93            label: Some("Pass Timing Resolve Buffer"),
94            size: buffer_size,
95            usage: wgpu::BufferUsages::QUERY_RESOLVE | wgpu::BufferUsages::COPY_SRC,
96            mapped_at_creation: false,
97        });
98        let slots = (0..READBACK_SLOTS)
99            .map(|_| ReadbackSlot {
100                buffer: device.create_buffer(&wgpu::BufferDescriptor {
101                    label: Some("Pass Timing Readback Buffer"),
102                    size: buffer_size,
103                    usage: wgpu::BufferUsages::COPY_DST | wgpu::BufferUsages::MAP_READ,
104                    mapped_at_creation: false,
105                }),
106                state: Arc::new(AtomicU8::new(SLOT_FREE)),
107                passes: RefCell::new(Vec::new()),
108            })
109            .collect();
110        Some(Self {
111            query_set,
112            resolve_buffer,
113            period_ns: queue.get_timestamp_period(),
114            cursor: Cell::new(0),
115            frame_passes: RefCell::new(Vec::new()),
116            labels: RefCell::new(Vec::new()),
117            totals: RefCell::new(Vec::new()),
118            slots,
119            frame_index: Cell::new(0),
120            frames_harvested: Cell::new(0),
121            span_nanoseconds: Cell::new(0),
122            dropped_passes: Cell::new(0),
123            dropped_frames: Cell::new(0),
124        })
125    }
126
127    pub(crate) fn query_set(&self) -> &wgpu::QuerySet {
128        &self.query_set
129    }
130
131    pub(crate) fn begin_pass(&self, label: &str) -> Option<(u32, u32)> {
132        let begin = self.cursor.get();
133        if begin + 2 > QUERY_CAPACITY {
134            self.dropped_passes.set(self.dropped_passes.get() + 1);
135            return None;
136        }
137        self.cursor.set(begin + 2);
138        let label_id = self.intern(label);
139        self.frame_passes.borrow_mut().push((label_id, begin));
140        Some((begin, begin + 1))
141    }
142
143    fn intern(&self, label: &str) -> u16 {
144        let mut labels = self.labels.borrow_mut();
145        if let Some(id) = labels.iter().position(|known| known == label) {
146            return id as u16;
147        }
148        labels.push(label.to_string());
149        self.totals.borrow_mut().push(LabelTotal::default());
150        (labels.len() - 1) as u16
151    }
152
153    pub(crate) fn harvest_completed(&self) {
154        for slot in &self.slots {
155            match slot.state.load(Ordering::Acquire) {
156                SLOT_MAPPED => {
157                    {
158                        let mapped = slot.buffer.slice(..).get_mapped_range();
159                        let span = accumulate_frame(
160                            &mut self.totals.borrow_mut(),
161                            &slot.passes.borrow(),
162                            &mapped,
163                            self.period_ns,
164                        );
165                        self.span_nanoseconds
166                            .set(self.span_nanoseconds.get().saturating_add(span));
167                    }
168                    slot.buffer.unmap();
169                    slot.passes.borrow_mut().clear();
170                    slot.state.store(SLOT_FREE, Ordering::Release);
171                    self.frames_harvested.set(self.frames_harvested.get() + 1);
172                }
173                SLOT_FAILED => {
174                    slot.passes.borrow_mut().clear();
175                    slot.state.store(SLOT_FREE, Ordering::Release);
176                }
177                _ => {}
178            }
179        }
180    }
181
182    pub(crate) fn frame_resolve(&self) -> Option<PendingResolve<'_>> {
183        let used = self.cursor.get();
184        if used == 0 {
185            return None;
186        }
187        let Some(slot_index) = self
188            .slots
189            .iter()
190            .position(|slot| slot.state.load(Ordering::Acquire) == SLOT_FREE)
191        else {
192            self.dropped_frames.set(self.dropped_frames.get() + 1);
193            return None;
194        };
195        Some(PendingResolve {
196            timer: self,
197            slot_index,
198            used,
199        })
200    }
201
202    pub(crate) fn finish_frame(&self) {
203        self.cursor.set(0);
204        self.frame_passes.borrow_mut().clear();
205        let frame = self.frame_index.get() + 1;
206        self.frame_index.set(frame);
207        if frame.is_multiple_of(PRINT_CADENCE_FRAMES) {
208            self.print_and_reset_window(frame);
209        }
210    }
211
212    pub(crate) fn report(&self) -> GpuPassTimingReport {
213        let labels = self.labels.borrow();
214        let totals = self.totals.borrow();
215        let mut entries: Vec<GpuPassTimingEntry> = labels
216            .iter()
217            .zip(totals.iter())
218            .filter(|(_, total)| total.passes > 0)
219            .map(|(label, total)| GpuPassTimingEntry {
220                label: label.clone(),
221                total_ms: total.nanoseconds as f64 / 1_000_000.0,
222                passes: total.passes,
223            })
224            .collect();
225        entries.sort_by(|a, b| b.total_ms.total_cmp(&a.total_ms));
226        GpuPassTimingReport {
227            frames: self.frames_harvested.get(),
228            span_ms: self.span_nanoseconds.get() as f64 / 1_000_000.0,
229            entries,
230        }
231    }
232
233    fn print_and_reset_window(&self, frame: u64) {
234        let report = self.report();
235        if report.frames > 0 {
236            let frames = f64::from(report.frames);
237            let total_ms: f64 = report.entries.iter().map(|entry| entry.total_ms).sum();
238            let mut line = format!(
239                "[GPU-PASS f#{frame}] frames={} span={:.2}ms/frame occupancy={:.2}ms/frame",
240                report.frames,
241                report.span_ms / frames,
242                total_ms / frames,
243            );
244            for entry in &report.entries {
245                line.push_str(&format!(
246                    " | {} {:.2}ms x{:.1}",
247                    entry.label,
248                    entry.total_ms / frames,
249                    entry.passes as f64 / frames,
250                ));
251            }
252            if self.dropped_passes.get() > 0 || self.dropped_frames.get() > 0 {
253                line.push_str(&format!(
254                    " | dropped: passes={} frames={}",
255                    self.dropped_passes.get(),
256                    self.dropped_frames.get(),
257                ));
258            }
259            eprintln!("{line}");
260        }
261        for total in self.totals.borrow_mut().iter_mut() {
262            *total = LabelTotal::default();
263        }
264        self.frames_harvested.set(0);
265        self.span_nanoseconds.set(0);
266        self.dropped_passes.set(0);
267        self.dropped_frames.set(0);
268    }
269}
270
271pub(crate) struct PendingResolve<'timer> {
272    timer: &'timer PassTimer,
273    slot_index: usize,
274    used: u32,
275}
276
277impl PendingResolve<'_> {
278    pub(crate) fn encode(&self, encoder: &mut wgpu::CommandEncoder) {
279        let slot = &self.timer.slots[self.slot_index];
280        encoder.resolve_query_set(
281            &self.timer.query_set,
282            0..self.used,
283            &self.timer.resolve_buffer,
284            0,
285        );
286        encoder.copy_buffer_to_buffer(
287            &self.timer.resolve_buffer,
288            0,
289            &slot.buffer,
290            0,
291            u64::from(self.used) * u64::from(wgpu::QUERY_SIZE),
292        );
293    }
294
295    pub(crate) fn arm_readback(self) {
296        let slot = &self.timer.slots[self.slot_index];
297        slot.passes
298            .borrow_mut()
299            .clone_from(&self.timer.frame_passes.borrow());
300        slot.state.store(SLOT_PENDING, Ordering::Release);
301        let state = Arc::clone(&slot.state);
302        slot.buffer
303            .slice(..)
304            .map_async(wgpu::MapMode::Read, move |result| {
305                let outcome = if result.is_ok() {
306                    SLOT_MAPPED
307                } else {
308                    SLOT_FAILED
309                };
310                state.store(outcome, Ordering::Release);
311            });
312    }
313}
314
315pub(crate) fn begin_timed_render_pass<'encoder>(
316    pass_timer: Option<&PassTimer>,
317    encoder: &'encoder mut wgpu::CommandEncoder,
318    descriptor: &wgpu::RenderPassDescriptor<'_>,
319) -> wgpu::RenderPass<'encoder> {
320    crate::frame_graph::note_render_pass(descriptor);
321    let timing = pass_timer.and_then(|timer| {
322        timer
323            .begin_pass(descriptor.label.unwrap_or("<unlabeled pass>"))
324            .map(|(begin, end)| (timer, begin, end))
325    });
326    let timestamp_writes = timing.map(|(timer, begin, end)| wgpu::RenderPassTimestampWrites {
327        query_set: timer.query_set(),
328        beginning_of_pass_write_index: Some(begin),
329        end_of_pass_write_index: Some(end),
330    });
331    encoder.begin_render_pass(&wgpu::RenderPassDescriptor {
332        timestamp_writes,
333        ..descriptor.clone()
334    })
335}
336
337fn accumulate_frame(
338    totals: &mut [LabelTotal],
339    passes: &[(u16, u32)],
340    mapped: &[u8],
341    period_ns: f32,
342) -> u64 {
343    let read_tick = |index: u32| -> Option<u64> {
344        let offset = index as usize * 8;
345        let bytes = mapped.get(offset..offset + 8)?;
346        Some(u64::from_le_bytes(bytes.try_into().expect("8-byte slice")))
347    };
348    let mut first_begin = u64::MAX;
349    let mut last_end = 0u64;
350    for &(label_id, begin_index) in passes {
351        let Some(total) = totals.get_mut(usize::from(label_id)) else {
352            continue;
353        };
354        let (Some(begin), Some(end)) = (read_tick(begin_index), read_tick(begin_index + 1)) else {
355            continue;
356        };
357        if end < begin {
358            continue;
359        }
360        first_begin = first_begin.min(begin);
361        last_end = last_end.max(end);
362        total.nanoseconds = total
363            .nanoseconds
364            .saturating_add(((end - begin) as f64 * f64::from(period_ns)) as u64);
365        total.passes += 1;
366    }
367    if last_end <= first_begin {
368        return 0;
369    }
370    ((last_end - first_begin) as f64 * f64::from(period_ns)) as u64
371}
372
373#[cfg(test)]
374mod tests {
375    use super::*;
376
377    fn ticks(values: &[u64]) -> Vec<u8> {
378        values.iter().flat_map(|v| v.to_le_bytes()).collect()
379    }
380
381    #[test]
382    fn accumulate_frame_attributes_ticks_by_label() {
383        let mut totals = vec![LabelTotal::default(); 2];
384        let mapped = ticks(&[1_000, 1_100, 1_100, 1_150, 1_150, 1_160]);
385        accumulate_frame(&mut totals, &[(0, 0), (1, 2), (0, 4)], &mapped, 2.0);
386        assert_eq!(totals[0].nanoseconds, 220);
387        assert_eq!(totals[0].passes, 2);
388        assert_eq!(totals[1].nanoseconds, 100);
389        assert_eq!(totals[1].passes, 1);
390    }
391
392    #[test]
393    fn accumulate_frame_skips_backwards_and_out_of_range_pairs() {
394        let mut totals = vec![LabelTotal::default(); 1];
395        let mapped = ticks(&[500, 400]);
396        accumulate_frame(&mut totals, &[(0, 0)], &mapped, 1.0);
397        assert_eq!(totals[0].passes, 0, "an end before its begin is skipped");
398
399        accumulate_frame(&mut totals, &[(0, 6)], &mapped, 1.0);
400        assert_eq!(totals[0].passes, 0, "indices past the mapping are skipped");
401
402        let mapped = ticks(&[100, 250]);
403        accumulate_frame(&mut totals, &[(9, 0)], &mapped, 1.0);
404        assert_eq!(totals[0].passes, 0, "an unknown label id is skipped");
405    }
406}