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