Skip to main content

vyre_driver_wgpu/engine/
persistent.rs

1//! Persistent-kernel queue execution for wgpu compute pipelines.
2//!
3//! This module provides host-side queue management and GPU dispatch for
4//! persistent kernels. A persistent pipeline is compiled once and reused
5//! across multiple work items, with only buffer contents changing between
6//! calls. If the wgpu device is lost, the backend recovers and the pipeline
7//! is recompiled automatically.
8
9use std::sync::Arc;
10
11use smallvec::SmallVec;
12use vyre_driver::BackendError;
13use vyre_driver::DispatchConfig;
14use vyre_driver::VyreBackend;
15use vyre_foundation::ir::Program;
16
17/// One unit of persistent-kernel work.
18#[derive(Clone, Debug, Eq, PartialEq)]
19pub struct PersistentPayloadWorkItem {
20    /// Stable work identifier.
21    pub id: u32,
22    /// Input payload consumed by the resident kernel.
23    pub payload: Vec<u8>,
24}
25
26/// Output produced for one persistent-kernel work item.
27#[derive(Clone, Debug, Eq, PartialEq)]
28pub struct WorkResult {
29    /// Stable work identifier copied from the input item.
30    pub id: u32,
31    /// Output payload produced by the kernel body.
32    pub payload: Vec<u8>,
33}
34
35/// GPU-side queue contract for persistent kernels.
36#[derive(Clone, Debug, Default, Eq, PartialEq)]
37pub struct PersistentQueue {
38    items: SmallVec<[PersistentPayloadWorkItem; 16]>,
39}
40
41impl PersistentQueue {
42    /// Create an empty persistent work queue.
43    #[must_use]
44    pub fn new() -> Self {
45        Self {
46            items: SmallVec::new(),
47        }
48    }
49
50    /// Enqueue one work item.
51    pub fn push(&mut self, item: PersistentPayloadWorkItem) {
52        self.items.push(item);
53    }
54
55    /// Number of queued work items.
56    #[must_use]
57    pub fn len(&self) -> usize {
58        self.items.len()
59    }
60
61    /// Returns true when the queue contains no work.
62    #[must_use]
63    pub fn is_empty(&self) -> bool {
64        self.items.is_empty()
65    }
66}
67
68/// Resident-kernel execution summary.
69#[derive(Clone, Debug, Eq, PartialEq)]
70pub struct PersistentKernelReport {
71    /// Number of kernel launches used to drain the queue.
72    pub kernel_launches: u32,
73    /// Results in queue order.
74    pub results: Vec<WorkResult>,
75}
76
77/// Compile a persistent pipeline and drain `queue` through the GPU.
78///
79/// The pipeline and bind-group layout are created once and reused for every
80/// work item. Only the input/output buffer contents change between launches.
81/// If the backend reports device loss before or during drainage, recovery is
82/// attempted and the pipeline is recompiled automatically.
83///
84/// # Errors
85///
86/// Returns [`BackendError`] when queue validation, pipeline compilation,
87/// device recovery, or GPU dispatch fails.
88pub fn run_persistent_kernel(
89    backend: &crate::WgpuBackend,
90    program: &Program,
91    config: &DispatchConfig,
92    queue: PersistentQueue,
93) -> Result<PersistentKernelReport, BackendError> {
94    if queue.is_empty() {
95        return Ok(PersistentKernelReport {
96            kernel_launches: 0,
97            results: Vec::new(),
98        });
99    }
100
101    let _work_items = u32::try_from(queue.items.len()).map_err(|_| {
102        BackendError::new(
103            "persistent queue length exceeds u32 GPU counters. Fix: shard work into multiple queues.",
104        )
105    })?;
106
107    let pipeline = ensure_persistent_pipeline(backend, program, config)?;
108    let mut payloads = SmallVec::<[&[u8]; 16]>::new();
109    payloads.try_reserve(queue.items.len()).map_err(|source| {
110        BackendError::new(format!(
111            "persistent kernel could not reserve {} payload slice reference(s): {source}. Fix: split the persistent queue before dispatch.",
112            queue.items.len()
113        ))
114    })?;
115    payloads.extend(queue.items.iter().map(|item| item.payload.as_slice()));
116    let output_batches = pipeline
117        .dispatch_coalesced_borrowed(&payloads, config)
118        .map_err(|error| {
119            BackendError::new(format!(
120                "persistent kernel coalesced dispatch failed for {} work items: {error}. Fix: verify the program and queued input payloads are compatible.",
121                queue.items.len()
122            ))
123        })?;
124    drop(payloads);
125    if output_batches.len() != queue.items.len() {
126        return Err(BackendError::new(format!(
127            "persistent kernel returned {} output batch(es) for {} queued work item(s). Fix: keep coalesced dispatch output cardinality identical to queue length.",
128            output_batches.len(),
129            queue.items.len()
130        )));
131    }
132
133    let mut results = Vec::new();
134    results.try_reserve(queue.items.len()).map_err(|source| {
135        BackendError::new(format!(
136            "persistent kernel could not reserve {} work result slot(s): {source}. Fix: split the persistent queue before collecting outputs.",
137            queue.items.len()
138        ))
139    })?;
140    for (item_index, (item, outputs)) in queue.items.into_iter().zip(output_batches).enumerate() {
141        if outputs.len() != 1 {
142            return Err(BackendError::new(format!(
143                "persistent kernel work item index {item_index} id={} returned {} output buffer(s); WorkResult requires exactly one payload. Fix: use a persistent program with one public output or extend PersistentKernelReport to carry multi-output results explicitly.",
144                item.id,
145                outputs.len()
146            )));
147        }
148        let mut outputs = outputs.into_iter();
149        let Some(payload) = outputs.next() else {
150            return Err(BackendError::new(format!(
151                "persistent kernel work item index {item_index} id={} returned no payload after output cardinality validation. Fix: keep persistent output extraction synchronized with the one-output WorkResult contract.",
152                item.id
153            )));
154        };
155        results.push(WorkResult {
156            id: item.id,
157            payload,
158        });
159    }
160
161    Ok(PersistentKernelReport {
162        kernel_launches: 1,
163        results,
164    })
165}
166
167fn ensure_persistent_pipeline(
168    backend: &crate::WgpuBackend,
169    program: &Program,
170    config: &DispatchConfig,
171) -> Result<Arc<crate::pipeline::WgpuPipeline>, BackendError> {
172    if backend.device_lost() {
173        return recover_and_recompile(backend, program, config);
174    }
175    backend.compile_persistent(program, config)
176}
177
178fn recover_and_recompile(
179    backend: &crate::WgpuBackend,
180    program: &Program,
181    config: &DispatchConfig,
182) -> Result<Arc<crate::pipeline::WgpuPipeline>, BackendError> {
183    backend.try_recover().map_err(|error| {
184        BackendError::new(format!(
185            "persistent kernel dispatch encountered device loss and recovery failed: {error}. Fix: ensure the GPU adapter is available before dispatching persistent work."
186        ))
187    })?;
188    backend.compile_persistent(program, config)
189}
190
191#[cfg(test)]
192mod tests {
193    use super::*;
194    use crate::WgpuBackend;
195    use vyre_driver::DispatchConfig;
196    use vyre_foundation::ir::{BufferDecl, DataType, Expr, Node, Program};
197
198    fn add_one_program(words: u32) -> Program {
199        let idx = Expr::gid_x();
200        let in_bounds = Expr::lt(idx.clone(), Expr::u32(words));
201        Program::wrapped(
202            vec![
203                BufferDecl::read("input", 0, DataType::U32).with_count(words),
204                BufferDecl::output("out", 1, DataType::U32)
205                    .with_count(words)
206                    .with_output_byte_range(0..(words as usize * 4)),
207            ],
208            [64, 1, 1],
209            vec![
210                Node::if_then(
211                    in_bounds,
212                    vec![Node::store(
213                        "out",
214                        idx.clone(),
215                        Expr::add(Expr::load("input", idx), Expr::u32(1)),
216                    )],
217                ),
218                Node::return_(),
219            ],
220        )
221    }
222
223    #[test]
224    fn persistent_kernel_queue_dispatches_on_gpu() {
225        let backend =
226            WgpuBackend::acquire().expect("Fix: GPU must be available for persistent kernel test");
227        let program = add_one_program(1);
228        let mut queue = PersistentQueue::new();
229        for id in 0..16 {
230            queue.push(PersistentPayloadWorkItem {
231                id,
232                payload: (id + 100).to_le_bytes().to_vec(),
233            });
234        }
235
236        let report = run_persistent_kernel(&backend, &program, &DispatchConfig::default(), queue)
237            .expect("Fix: persistent kernel dispatch must succeed");
238
239        assert_eq!(report.kernel_launches, 1);
240        assert_eq!(report.results.len(), 16);
241        for result in &report.results {
242            let input_val = result.id + 100;
243            let expected = input_val + 1;
244            let actual = u32::from_le_bytes([
245                result.payload[0],
246                result.payload[1],
247                result.payload[2],
248                result.payload[3],
249            ]);
250            assert_eq!(
251                actual, expected,
252                "work item {}: expected {}, got {}",
253                result.id, expected, actual
254            );
255        }
256    }
257
258    #[test]
259    fn persistent_kernel_survives_device_loss_recovery() {
260        let backend = WgpuBackend::acquire()
261            .expect("Fix: GPU must be available for persistent kernel recovery test");
262        let program = add_one_program(1);
263        let mut queue = PersistentQueue::new();
264        queue.push(PersistentPayloadWorkItem {
265            id: 0,
266            payload: 42u32.to_le_bytes().to_vec(),
267        });
268
269        // First dispatch.
270        let report1 = run_persistent_kernel(
271            &backend,
272            &program,
273            &DispatchConfig::default(),
274            queue.clone(),
275        )
276        .expect("Fix: first persistent dispatch must succeed");
277        assert_eq!(report1.results.len(), 1);
278        let actual1 = u32::from_le_bytes([
279            report1.results[0].payload[0],
280            report1.results[0].payload[1],
281            report1.results[0].payload[2],
282            report1.results[0].payload[3],
283        ]);
284        assert_eq!(actual1, 43);
285
286        // Simulate device loss.
287        backend
288            .force_device_lost()
289            .expect("Fix: test hook must invalidate the cached device");
290        assert!(backend.device_lost());
291
292        // Recovery should happen automatically inside run_persistent_kernel.
293        let report2 = run_persistent_kernel(&backend, &program, &DispatchConfig::default(), queue)
294            .expect("Fix: persistent dispatch after recovery must succeed");
295        assert_eq!(report2.results.len(), 1);
296        let actual2 = u32::from_le_bytes([
297            report2.results[0].payload[0],
298            report2.results[0].payload[1],
299            report2.results[0].payload[2],
300            report2.results[0].payload[3],
301        ]);
302        assert_eq!(actual2, 43);
303    }
304}