1use 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#[derive(Clone, Debug, Eq, PartialEq)]
19pub struct PersistentPayloadWorkItem {
20 pub id: u32,
22 pub payload: Vec<u8>,
24}
25
26#[derive(Clone, Debug, Eq, PartialEq)]
28pub struct WorkResult {
29 pub id: u32,
31 pub payload: Vec<u8>,
33}
34
35#[derive(Clone, Debug, Default, Eq, PartialEq)]
37pub struct PersistentQueue {
38 items: SmallVec<[PersistentPayloadWorkItem; 16]>,
39}
40
41impl PersistentQueue {
42 #[must_use]
44 pub fn new() -> Self {
45 Self {
46 items: SmallVec::new(),
47 }
48 }
49
50 pub fn push(&mut self, item: PersistentPayloadWorkItem) {
52 self.items.push(item);
53 }
54
55 #[must_use]
57 pub fn len(&self) -> usize {
58 self.items.len()
59 }
60
61 #[must_use]
63 pub fn is_empty(&self) -> bool {
64 self.items.is_empty()
65 }
66}
67
68#[derive(Clone, Debug, Eq, PartialEq)]
70pub struct PersistentKernelReport {
71 pub kernel_launches: u32,
73 pub results: Vec<WorkResult>,
75}
76
77pub 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 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 backend
288 .force_device_lost()
289 .expect("Fix: test hook must invalidate the cached device");
290 assert!(backend.device_lost());
291
292 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}