dynamis-gpu 0.7.0

wgpu device, buffer, stream, pipeline, and readback runtime
Documentation
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()
    }
}