use crate::{Readback, SubmissionEncoder};
use std::collections::VecDeque;
use wgpu::{
Buffer, BufferDescriptor, BufferUsages, ComputePassTimestampWrites, Device, QUERY_SIZE,
QuerySet, QuerySetDescriptor, QueryType,
};
const QUERY_BYTES: usize = QUERY_SIZE as usize;
#[derive(Clone, Debug, PartialEq)]
pub struct GpuPassTiming {
pub label: &'static str,
pub nanoseconds: f64,
}
impl GpuPassTiming {
pub fn milliseconds(&self) -> f64 {
self.nanoseconds / 1e6
}
}
pub struct GpuTimer {
query_set: QuerySet,
resolved: Buffer,
readback: Readback,
labels: Vec<&'static str>,
period_ns: f32,
ran: VecDeque<Vec<bool>>,
sequence: u64,
}
impl GpuTimer {
pub fn new(
device: &Device,
labels: &[&'static str],
period_ns: f32,
label_prefix: &str,
) -> Self {
assert!(!labels.is_empty(), "a gpu timer needs at least one pass");
assert!(
period_ns > 0.0,
"timestamp period must be strictly positive"
);
let resolved_bytes = Self::resolved_bytes(labels);
let query_set = device.create_query_set(&QuerySetDescriptor {
label: Some(&format!("{label_prefix} timestamps")),
ty: QueryType::Timestamp,
count: (labels.len() * 2) as u32,
});
let resolved = device.create_buffer(&BufferDescriptor {
label: Some(&format!("{label_prefix} resolved timestamps")),
size: resolved_bytes,
usage: BufferUsages::QUERY_RESOLVE | BufferUsages::COPY_SRC,
mapped_at_creation: false,
});
let readback = Readback::new(
device,
&format!("{label_prefix} timestamp readback"),
resolved_bytes,
Readback::DEPTH,
);
Self {
query_set,
resolved,
readback,
labels: labels.to_vec(),
period_ns,
ran: VecDeque::new(),
sequence: 0,
}
}
pub fn pass_count(&self) -> usize {
self.labels.len()
}
fn resolved_bytes(labels: &[&'static str]) -> u64 {
(labels.len() * 2 * QUERY_BYTES) as u64
}
pub fn writes(&self, slot: usize) -> ComputePassTimestampWrites<'_> {
assert!(
slot < self.labels.len(),
"timing slot {slot} exceeds the {n} recorded passes",
n = self.labels.len()
);
ComputePassTimestampWrites {
query_set: &self.query_set,
beginning_of_pass_write_index: Some((slot * 2) as u32),
end_of_pass_write_index: Some((slot * 2 + 1) as u32),
}
}
pub fn capture(
&mut self,
encoder: &mut SubmissionEncoder,
ran: &[bool],
) -> Option<Vec<GpuPassTiming>> {
assert_eq!(
ran.len(),
self.labels.len(),
"a captured frame declares one activity flag per pass"
);
let count = (self.labels.len() * 2) as u32;
encoder.resolve_query_set(&self.query_set, 0..count, &self.resolved, 0);
let displaced = self.readback.enqueue(
encoder,
&self.resolved,
0,
self.resolved.size(),
self.sequence,
);
self.sequence += 1;
self.ran.push_back(ran.to_vec());
displaced.map(|(_, bytes)| {
let ran = self.retired();
self.decode(&bytes, &ran)
})
}
pub fn collect(&mut self) -> Vec<Vec<GpuPassTiming>> {
let retired = self.readback.collect();
retired
.into_iter()
.map(|(_, bytes)| {
let ran = self.retired();
self.decode(&bytes, &ran)
})
.collect()
}
fn retired(&mut self) -> Vec<bool> {
self.ran
.pop_front()
.expect("a retired frame declares the passes that ran")
}
fn decode(&self, bytes: &[u8], ran: &[bool]) -> Vec<GpuPassTiming> {
let (chunks, remainder) = bytes.as_chunks::<QUERY_BYTES>();
assert!(
remainder.is_empty(),
"resolved timestamp buffer is not a multiple of {QUERY_BYTES} bytes"
);
assert_eq!(
chunks.len(),
self.labels.len() * 2,
"resolved timestamp count disagrees with the recorded passes"
);
let ticks: Vec<u64> = chunks
.iter()
.map(|chunk| u64::from_le_bytes(*chunk))
.collect();
self.labels
.iter()
.enumerate()
.filter(|(slot, _)| ran[*slot])
.map(|(slot, label)| GpuPassTiming {
label,
nanoseconds: ticks[slot * 2 + 1].wrapping_sub(ticks[slot * 2]) as f64
* self.period_ns as f64,
})
.collect()
}
}