Skip to main content

dynamis_gpu/
readback.rs

1use crate::SubmissionEncoder;
2use std::collections::VecDeque;
3use std::sync::mpsc::{self, Receiver, TryRecvError};
4use std::sync::{Arc, OnceLock};
5use std::time::{Duration, Instant};
6use wgpu::{
7    Buffer, BufferAddress, BufferAsyncError, BufferDescriptor, BufferUsages, Device, MapMode,
8    PollType, Queue, SubmissionIndex,
9};
10
11const READBACK_TIMEOUT: Duration = Duration::from_secs(30);
12
13struct Pending {
14    sequence: u64,
15    bytes: BufferAddress,
16    submission: Arc<OnceLock<SubmissionIndex>>,
17    completion: Receiver<Result<(), BufferAsyncError>>,
18}
19
20struct Slot {
21    label: String,
22    staging: Buffer,
23    pending: Option<Pending>,
24}
25
26impl Slot {
27    fn new(device: &Device, label: &str, size: BufferAddress) -> Self {
28        Self {
29            label: label.to_owned(),
30            staging: device.create_buffer(&BufferDescriptor {
31                label: Some(label),
32                size,
33                usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
34                mapped_at_creation: false,
35            }),
36            pending: None,
37        }
38    }
39
40    fn is_sealed(&self) -> bool {
41        self.pending
42            .as_ref()
43            .is_some_and(|pending| pending.submission.get().is_some())
44    }
45
46    fn record(
47        &mut self,
48        encoder: &mut SubmissionEncoder,
49        regions: &[(&Buffer, BufferAddress, BufferAddress)],
50        sequence: u64,
51    ) {
52        assert!(
53            self.pending.is_none(),
54            "readback {} is still in use",
55            self.label
56        );
57        let mut at = 0;
58        for (source, source_offset, bytes) in regions {
59            assert!(
60                *bytes > 0 && bytes.is_multiple_of(4),
61                "readback length must be positive and word aligned"
62            );
63            assert!(
64                source_offset.is_multiple_of(4),
65                "readback source offset must be word aligned"
66            );
67            assert!(
68                *source_offset <= source.size() && *bytes <= source.size() - source_offset,
69                "readback exceeds source buffer"
70            );
71            encoder.copy_buffer_to_buffer(source, *source_offset, &self.staging, at, *bytes);
72            at += bytes;
73        }
74        assert!(
75            at > 0 && at <= self.staging.size(),
76            "readback exceeds staging capacity"
77        );
78        let (sender, completion) = mpsc::channel();
79        encoder.map_buffer_on_submit(&self.staging, MapMode::Read, ..at, move |result| {
80            let _ = sender.send(result);
81        });
82        self.pending = Some(Pending {
83            sequence,
84            bytes: at,
85            submission: encoder.submission(),
86            completion,
87        });
88    }
89
90    fn collect(&mut self) -> Option<(u64, Vec<u8>)> {
91        let pending = self.pending.as_ref()?;
92        pending.submission.get()?;
93        match pending.completion.try_recv() {
94            Ok(result) => Some(self.consume(result)),
95            Err(TryRecvError::Empty) => None,
96            Err(TryRecvError::Disconnected) => {
97                panic!("readback {} completion was dropped", self.label)
98            }
99        }
100    }
101
102    fn wait(&mut self, device: &Device) -> (u64, Vec<u8>) {
103        let pending = self
104            .pending
105            .as_ref()
106            .expect("cannot wait for an idle readback");
107        let submission = pending.submission.get().cloned().unwrap_or_else(|| {
108            panic!(
109                "readback {} must be submitted before waiting or reusing its slot",
110                self.label
111            )
112        });
113        let deadline = Instant::now() + READBACK_TIMEOUT;
114        device
115            .poll(PollType::Wait {
116                submission_index: Some(submission),
117                timeout: Some(READBACK_TIMEOUT),
118            })
119            .unwrap_or_else(|error| panic!("readback {} submission failed: {error}", self.label));
120        let result = pending
121            .completion
122            .recv_timeout(deadline.saturating_duration_since(Instant::now()))
123            .unwrap_or_else(|error| panic!("readback {} completion failed: {error}", self.label));
124        self.consume(result)
125    }
126
127    fn consume(&mut self, result: Result<(), BufferAsyncError>) -> (u64, Vec<u8>) {
128        result.unwrap_or_else(|error| panic!("readback {} mapping failed: {error}", self.label));
129        let pending = self
130            .pending
131            .take()
132            .expect("readback completion requires a pending read");
133        let bytes = self
134            .staging
135            .slice(..pending.bytes)
136            .get_mapped_range()
137            .expect("completed readback must be mapped")
138            .to_vec();
139        self.staging.unmap();
140        (pending.sequence, bytes)
141    }
142}
143
144pub const EVENT_SLOTS: u32 = Readback::DEPTH as u32 + 2;
145
146pub struct Readback {
147    device: Device,
148    label: String,
149    size: BufferAddress,
150    slots: VecDeque<Slot>,
151    inflight: usize,
152    last_sequence: Option<u64>,
153}
154
155impl Readback {
156    pub const DEPTH: usize = 4;
157
158    pub fn new(device: &Device, label: &str, size: BufferAddress, depth: usize) -> Self {
159        assert!(
160            size > 0 && size.is_multiple_of(4),
161            "readback size must be positive and word aligned"
162        );
163        assert!(depth > 0, "a readback needs at least one staging slot");
164        Self {
165            device: device.clone(),
166            label: label.to_owned(),
167            size,
168            slots: (0..depth)
169                .map(|index| Slot::new(device, &format!("{label} slot {index}"), size))
170                .collect(),
171            inflight: 0,
172            last_sequence: None,
173        }
174    }
175
176    pub fn size(&self) -> BufferAddress {
177        self.size
178    }
179
180    pub fn is_idle(&self) -> bool {
181        self.inflight == 0
182    }
183
184    pub fn enqueue(
185        &mut self,
186        encoder: &mut SubmissionEncoder,
187        source: &Buffer,
188        source_offset: BufferAddress,
189        bytes: BufferAddress,
190        sequence: u64,
191    ) -> Option<(u64, Vec<u8>)> {
192        self.enqueue_regions(encoder, &[(source, source_offset, bytes)], sequence)
193    }
194
195    pub fn enqueue_regions(
196        &mut self,
197        encoder: &mut SubmissionEncoder,
198        regions: &[(&Buffer, BufferAddress, BufferAddress)],
199        sequence: u64,
200    ) -> Option<(u64, Vec<u8>)> {
201        assert!(
202            self.last_sequence
203                .is_none_or(|previous| sequence > previous),
204            "readback {} sequences must increase",
205            self.label
206        );
207        let displaced = self.reclaim();
208        self.slots[self.inflight].record(encoder, regions, sequence);
209        self.inflight += 1;
210        self.last_sequence = Some(sequence);
211        displaced
212    }
213
214    pub fn collect(&mut self) -> Vec<(u64, Vec<u8>)> {
215        let mut completed = Vec::new();
216        while self.inflight > 0 {
217            let Some(entry) = self.slots[0].collect() else {
218                break;
219            };
220            self.retire();
221            completed.push(entry);
222        }
223        completed
224    }
225
226    pub fn drain(&mut self) -> Vec<(u64, Vec<u8>)> {
227        let mut completed = Vec::with_capacity(self.inflight);
228        while self.inflight > 0 {
229            let entry = self.slots[0].wait(&self.device);
230            self.retire();
231            completed.push(entry);
232        }
233        completed
234    }
235
236    fn reclaim(&mut self) -> Option<(u64, Vec<u8>)> {
237        if self.inflight < self.slots.len() {
238            return None;
239        }
240        if !self.slots[0].is_sealed() {
241            let index = self.slots.len();
242            self.slots.push_back(Slot::new(
243                &self.device,
244                &format!("{} slot {index}", self.label),
245                self.size,
246            ));
247            return None;
248        }
249        let mut oldest = self.slots.pop_front().expect("readback holds a slot");
250        let entry = oldest.wait(&self.device);
251        self.slots.push_back(oldest);
252        self.inflight -= 1;
253        Some(entry)
254    }
255
256    fn retire(&mut self) {
257        let slot = self.slots.pop_front().expect("readback holds a slot");
258        self.slots.push_back(slot);
259        self.inflight -= 1;
260        if self.inflight == 0 {
261            self.last_sequence = None;
262        }
263    }
264}
265
266pub fn read_regions(
267    device: &Device,
268    queue: &Queue,
269    label: &str,
270    regions: &[(&Buffer, BufferAddress, BufferAddress)],
271) -> Vec<u8> {
272    let bytes: BufferAddress = regions.iter().map(|region| region.2).sum();
273    let mut readback = Readback::new(device, label, bytes, 1);
274    let mut encoder = SubmissionEncoder::new(device, label);
275    assert!(
276        readback.enqueue_regions(&mut encoder, regions, 0).is_none(),
277        "a one shot read requires an idle readback"
278    );
279    encoder.submit(queue);
280    readback
281        .drain()
282        .pop()
283        .expect("a one shot read retires exactly once")
284        .1
285}