1use crate::staging_reserve::{reserve_backend_vec, reserve_smallvec, reserve_vec};
4use crate::{AdapterRecoveryTarget, DispatchArena, WgpuBackend};
5use std::hash::{BuildHasherDefault, Hasher};
6use std::sync::{
7 atomic::{AtomicBool, Ordering},
8 Arc,
9};
10use std::time::Instant;
11use vyre_driver::persistent::PersistentThreadMode;
12use vyre_driver::speculate::SpeculationMode;
13use vyre_foundation::ir::Program;
14
15fn empty_batch_result_slots<T>(
16 len: usize,
17) -> Result<Vec<Option<Result<T, vyre_driver::BackendError>>>, vyre_driver::BackendError> {
18 let mut slots = Vec::new();
19 reserve_vec(
20 &mut slots,
21 len,
22 "WGPU backend",
23 "batch result slot",
24 "split the batch before dispatch",
25 )?;
26 slots.resize_with(len, || None);
27 Ok(slots)
28}
29
30fn finalize_batch_results<T>(
31 slots: Vec<Option<Result<T, vyre_driver::BackendError>>>,
32 missing_slot_message: &'static str,
33) -> Result<Vec<Result<T, vyre_driver::BackendError>>, vyre_driver::BackendError> {
34 let mut results = Vec::new();
35 reserve_vec(
36 &mut results,
37 slots.len(),
38 "WGPU backend",
39 "final batch result",
40 "split the batch before dispatch",
41 )?;
42 for slot in slots {
43 results.push(
44 slot.unwrap_or_else(|| Err(vyre_driver::BackendError::new(missing_slot_message))),
45 );
46 }
47 Ok(results)
48}
49
50fn elapsed_micros_u64(start: Instant, label: &str) -> Result<u64, vyre_driver::BackendError> {
51 u64::try_from(start.elapsed().as_micros()).map_err(|source| {
52 vyre_driver::BackendError::new(format!(
53 "{label} elapsed time cannot fit u64 microseconds: {source}. Fix: split or timeout the dispatch before telemetry overflows."
54 ))
55 })
56}
57
58impl WgpuBackend {
59 #[must_use]
61 pub fn adapter_info(&self) -> &wgpu::AdapterInfo {
62 &self.adapter_info
63 }
64
65 #[must_use]
67 pub fn device_limits(&self) -> &wgpu::Limits {
68 &self.device_limits
69 }
70
71 pub fn acquire() -> Result<Self, vyre_driver::BackendError> {
74 let ((device, queue), adapter_info, enabled_features) = crate::runtime::init_device()
75 .map_err(|error| {
76 let report = crate::runtime::device::adapter_probe_report();
77 vyre_driver::BackendError::new(format!(
78 "no compatible GPU adapter found. Probed adapters: [{}]. Missing features / limits: [{}]. Underlying error: {error}. Fix: install a compatible GPU driver and ensure a wgpu-supported backend (Vulkan, Metal, DX12) is available.",
79 report.probed.join(", "),
80 if report.missing.is_empty() {
81 "none".to_string()
82 } else {
83 report.missing.join(", ")
84 }
85 ))
86 })?;
87 let recovery_target = AdapterRecoveryTarget::Identity(
88 crate::runtime::device::AdapterIdentity::from_info(&adapter_info),
89 );
90 Self::from_device_queue(
91 device,
92 queue,
93 adapter_info,
94 enabled_features,
95 recovery_target,
96 )
97 }
98
99 pub fn acquire_adapter(index: usize) -> Result<Self, vyre_driver::BackendError> {
101 let ((device, queue), adapter_info, enabled_features) =
102 crate::runtime::device::init_device_for_adapter(index)
103 .map_err(|error| vyre_driver::BackendError::new(error.to_string()))?;
104 Self::from_device_queue(
105 device,
106 queue,
107 adapter_info,
108 enabled_features,
109 AdapterRecoveryTarget::Index(index),
110 )
111 }
112
113 fn from_device_queue(
114 device: wgpu::Device,
115 queue: wgpu::Queue,
116 adapter_info: wgpu::AdapterInfo,
117 enabled_features: crate::runtime::device::EnabledFeatures,
118 recovery_target: AdapterRecoveryTarget,
119 ) -> Result<Self, vyre_driver::BackendError> {
120 let device_limits = device.limits();
121 let adapter_name = Arc::<str>::from(adapter_info.name.as_str());
122 let cache_tiers = vec![
123 crate::runtime::cache::CacheTier::try_new("hot", 1 << 24)?,
124 crate::runtime::cache::CacheTier::try_new("cold", 1 << 30)?,
125 ];
126 let persistent_pool = crate::buffer::BufferPool::with_tiering(
127 device.clone(),
128 queue.clone(),
129 &vyre_driver::DispatchConfig::default(),
130 cache_tiers,
131 )?;
132 let (pipeline_cache_entries, pipeline_cache_bytes) =
133 vyre_driver::pipeline::pipeline_cache_limits_from_env();
134 Ok(Self {
135 adapter_name,
136 adapter_info,
137 device_limits,
138 device_queue: Arc::new(arc_swap::ArcSwap::new(Arc::new((
139 device.clone(),
140 queue.clone(),
141 )))),
142 dispatch_arena: Arc::new(arc_swap::ArcSwap::from_pointee(DispatchArena::new(
143 device.clone(),
144 queue.clone(),
145 &vyre_driver::DispatchConfig::default(),
146 ))),
147 persistent_pool: Arc::new(arc_swap::ArcSwap::new(Arc::new(persistent_pool))),
148 pipeline_cache: Arc::new(
149 crate::runtime::cache::pipeline::LruPipelineCache::with_limits(
150 pipeline_cache_entries,
151 pipeline_cache_bytes,
152 ),
153 ),
154 wgsl_dispatch_pipeline_cache: Arc::new(dashmap::DashMap::with_hasher(
155 BuildHasherDefault::<rustc_hash::FxHasher>::default(),
156 )),
157 resident_pipeline_cache: Arc::new(dashmap::DashMap::with_hasher(BuildHasherDefault::<
158 rustc_hash::FxHasher,
159 >::default(
160 ))),
161 validation_cache: Arc::new(vyre_driver::validation::ValidationCache::default()),
162 shape_history: Arc::new(std::sync::Mutex::new(
163 vyre_driver::shape_prediction::ShapeHistory::new(),
164 )),
165 predicted_programs: Arc::new(dashmap::DashMap::with_hasher(BuildHasherDefault::<
166 rustc_hash::FxHasher,
167 >::default())),
168 bind_group_layout_cache: Arc::new(dashmap::DashMap::with_hasher(BuildHasherDefault::<
169 rustc_hash::FxHasher,
170 >::default(
171 ))),
172 resident_handles: Arc::new(dashmap::DashMap::with_hasher(BuildHasherDefault::<
173 rustc_hash::FxHasher,
174 >::default())),
175 device_lost: Arc::new(AtomicBool::new(false)),
176 enabled_features,
177 recovery_target,
178 })
179 }
180
181 pub(crate) fn current_device_queue(&self) -> Arc<(wgpu::Device, wgpu::Queue)> {
182 self.device_queue.load_full()
183 }
184
185 #[must_use]
187 pub fn device_queue(&self) -> Arc<(wgpu::Device, wgpu::Queue)> {
188 self.current_device_queue()
189 }
190
191 pub(crate) fn current_persistent_pool(&self) -> crate::buffer::BufferPool {
192 self.persistent_pool.load_full().as_ref().clone()
193 }
194
195 fn resident_pipeline_cache_key(
196 &self,
197 program: &Program,
198 config: &vyre_driver::DispatchConfig,
199 ) -> Result<(u64, u64, usize), vyre_driver::BackendError> {
200 let wire = program.to_wire().map_err(|source| {
201 vyre_driver::BackendError::new(format!(
202 "WGPU resident pipeline cache could not encode Program: {source}. Fix: validate the Program before resident dispatch."
203 ))
204 })?;
205 let mut program_hasher = rustc_hash::FxHasher::default();
206 program_hasher.write(&wire);
207 let mut config_hasher = rustc_hash::FxHasher::default();
208 config_hasher.write(format!("{config:?}").as_bytes());
209 Ok((program_hasher.finish(), config_hasher.finish(), wire.len()))
210 }
211
212 pub(crate) fn compile_resident_pipeline_cached(
213 &self,
214 program: &Program,
215 config: &vyre_driver::DispatchConfig,
216 ) -> Result<Arc<crate::pipeline::WgpuPipeline>, vyre_driver::BackendError> {
217 let key = self.resident_pipeline_cache_key(program, config)?;
218 if let Some(hit) = self.resident_pipeline_cache.get(&key) {
219 return Ok(hit.clone());
220 }
221 self.enforce_config_caps(config)?;
222 self.validate_with_cache(program)?;
223 let compiled = crate::pipeline::WgpuPipeline::compile_with_device_queue(
224 program,
225 config,
226 self.adapter_info.clone(),
227 self.enabled_features,
228 self.current_device_queue(),
229 self.dispatch_arena_snapshot(),
230 self.current_persistent_pool(),
231 self.pipeline_cache.clone(),
232 self.bind_group_layout_cache.clone(),
233 )?;
234 match self.resident_pipeline_cache.entry(key) {
235 dashmap::mapref::entry::Entry::Occupied(entry) => Ok(entry.get().clone()),
236 dashmap::mapref::entry::Entry::Vacant(entry) => {
237 entry.insert(compiled.clone());
238 Ok(compiled)
239 }
240 }
241 }
242
243 pub(crate) fn dispatch_arena_snapshot(&self) -> Arc<DispatchArena> {
244 self.dispatch_arena.load_full()
245 }
246
247 pub(crate) fn validate_with_cache(
248 &self,
249 program: &Program,
250 ) -> Result<(), vyre_driver::BackendError> {
251 self.validation_cache.get_or_validate_backend(program, self)
252 }
253
254 pub fn force_device_lost(&self) -> Result<(), vyre_driver::BackendError> {
257 self.device_lost.store(true, Ordering::Release);
258 self.pipeline_cache.clear();
259 self.wgsl_dispatch_pipeline_cache.clear();
260 self.bind_group_layout_cache.clear();
261 self.validation_cache.clear()?;
262 let device_queue = self.device_queue.load_full();
263 self.dispatch_arena.store(Arc::new(DispatchArena::new(
264 device_queue.0.clone(),
265 device_queue.1.clone(),
266 &vyre_driver::DispatchConfig::default(),
267 )));
268 Ok(())
269 }
270
271 pub fn invalidate_impacted_pipeline_cache(
273 &self,
274 intervention_mask: &[u32],
275 rule_adj: &[u32],
276 state: &[u32],
277 join_rules: &[u32],
278 n: u32,
279 max_iterations: u32,
280 pipeline_lineage_cell: &[u32],
281 pipeline_keys: &[[u8; 32]],
282 ) -> Result<(), vyre_driver::BackendError> {
283 let final_impact_mask = vyre_driver::cache_invalidation::impacted_entries(
284 self,
285 intervention_mask,
286 rule_adj,
287 state,
288 join_rules,
289 n,
290 max_iterations,
291 pipeline_lineage_cell,
292 )
293 .map_err(|error| vyre_driver::BackendError::new(error.to_string()))?;
294 self.pipeline_cache
295 .invalidate_impacted(&final_impact_mask, pipeline_keys);
296 Ok(())
297 }
298
299 pub fn invalidate_pipeline_cache_for_changed_op(
301 &self,
302 changed_op_handle: u32,
303 pipeline_lineage_cell: &[u32],
304 pipeline_keys: &[[u8; 32]],
305 ) -> Result<(), vyre_driver::BackendError> {
306 let n = 1u32;
307 let rule_adj = vec![1u32];
308 let intervention_mask = vec![1u32];
309 let state = vec![1u32];
310 let join_rules = vec![1u32];
311 let max_iterations = 1u32;
312 let mut normalized_lineage_cell = Vec::with_capacity(pipeline_lineage_cell.len());
313 normalized_lineage_cell.extend(pipeline_lineage_cell.iter().map(|&op| {
314 if op == changed_op_handle {
315 0
316 } else {
317 u32::MAX
318 }
319 }));
320 self.invalidate_impacted_pipeline_cache(
321 &intervention_mask,
322 &rule_adj,
323 &state,
324 &join_rules,
325 n,
326 max_iterations,
327 &normalized_lineage_cell,
328 pipeline_keys,
329 )
330 }
331
332 pub fn invalidate_impacted_disk_cache(
334 &self,
335 intervention_mask: &[u32],
336 rule_adj: &[u32],
337 state: &[u32],
338 join_rules: &[u32],
339 n: u32,
340 max_iterations: u32,
341 pipeline_lineage_cell: &[u32],
342 cache_keys: &[String],
343 ) -> Result<(), vyre_driver::BackendError> {
344 crate::pipeline::disk_cache::invalidate_impacted(
345 self,
346 intervention_mask,
347 rule_adj,
348 state,
349 join_rules,
350 n,
351 max_iterations,
352 pipeline_lineage_cell,
353 cache_keys,
354 )
355 .map_err(|e| vyre_driver::BackendError::new(e.to_string()))
356 }
357
358 #[must_use]
360 #[inline]
361 pub fn new() -> Result<Self, vyre_driver::BackendError> {
362 Self::acquire().map_err(|e| vyre_driver::BackendError::new(e.to_string()))
363 }
364
365 pub fn shared() -> Result<Arc<Self>, vyre_driver::BackendError> {
367 static SHARED: std::sync::OnceLock<Result<Arc<WgpuBackend>, String>> =
368 std::sync::OnceLock::new();
369 match SHARED.get_or_init(|| Self::new().map(Arc::new).map_err(|e| e.to_string())) {
370 Ok(arc) => Ok(arc.clone()),
371 Err(msg) => Err(vyre_driver::BackendError::new(msg.clone())),
372 }
373 }
374
375 pub fn dispatch_borrowed_for_each_mapped_output<F>(
377 &self,
378 program: &Program,
379 inputs: &[&[u8]],
380 config: &vyre_driver::DispatchConfig,
381 visitor: F,
382 ) -> Result<(), vyre_driver::BackendError>
383 where
384 F: FnMut(usize, &[u8]) -> Result<(), vyre_driver::BackendError>,
385 {
386 let _span = tracing::trace_span!(
387 "vyre.dispatch_mapped_outputs",
388 backend = "wgpu",
389 inputs = inputs.len(),
390 label = tracing::field::Empty,
391 );
392 let _enter = _span.enter();
393 if let Some(label) = config.label.as_deref() {
394 _span.record("label", label);
395 }
396 let start = Instant::now();
397 self.dispatch_borrowed_async(program, inputs, config)?
398 .await_mapped_outputs(visitor)?;
399 tracing::trace!(
400 target: "vyre.dispatch",
401 elapsed_us = elapsed_micros_u64(start, "mapped-output dispatch")?,
402 inputs = inputs.len(),
403 "mapped-output dispatch completed"
404 );
405 Ok(())
406 }
407
408 pub fn dispatch_borrowed_for_each_pod_output<T, F>(
410 &self,
411 program: &Program,
412 inputs: &[&[u8]],
413 config: &vyre_driver::DispatchConfig,
414 mut visitor: F,
415 ) -> Result<(), vyre_driver::BackendError>
416 where
417 T: bytemuck::Pod,
418 F: FnMut(usize, &[T]) -> Result<(), vyre_driver::BackendError>,
419 {
420 self.dispatch_borrowed_for_each_mapped_output(program, inputs, config, |index, bytes| {
421 let typed = bytemuck::try_cast_slice::<u8, T>(bytes).map_err(|error| {
422 vyre_driver::BackendError::new(format!(
423 "mapped output #{index} cannot be viewed as {}: {error}. Fix: set output_byte_range to a length and offset aligned for the requested POD type.",
424 std::any::type_name::<T>()
425 ))
426 })?;
427 visitor(index, typed)
428 })
429 }
430
431 pub(crate) fn enforce_config_caps(
433 &self,
434 config: &vyre_driver::DispatchConfig,
435 ) -> Result<(), vyre_driver::BackendError> {
436 if matches!(config.speculation, Some(SpeculationMode::Force))
437 && !<Self as vyre_driver::VyreBackend>::supports_speculation(self)
438 {
439 return Err(vyre_driver::BackendError::UnsupportedFeature {
440 name: "speculative dispatch".to_string(),
441 backend: <Self as vyre_driver::VyreBackend>::id(self).to_string(),
442 });
443 }
444 if matches!(config.persistent_thread, Some(PersistentThreadMode::Force))
445 && !<Self as vyre_driver::VyreBackend>::supports_persistent_thread_dispatch(self)
446 {
447 return Err(vyre_driver::BackendError::UnsupportedFeature {
448 name: "persistent-thread dispatch".to_string(),
449 backend: <Self as vyre_driver::VyreBackend>::id(self).to_string(),
450 });
451 }
452 Ok(())
453 }
454
455 pub fn dispatch_speculative_prefilter_confirm<F>(
457 &self,
458 speculator: &vyre_driver::speculate::AdaptiveSpeculator,
459 plan: vyre_driver::speculate::SpeculativeDispatchPlan<'_>,
460 inputs: &[&[u8]],
461 config: &vyre_driver::DispatchConfig,
462 confirm_serial: F,
463 ) -> Result<vyre_driver::speculate::SpeculativeDispatchOutcome, vyre_driver::BackendError>
464 where
465 F: FnMut(
466 vyre_driver::OutputBuffers,
467 ) -> Result<vyre_driver::OutputBuffers, vyre_driver::BackendError>,
468 {
469 vyre_driver::speculate::dispatch_prefilter_confirm(
470 self,
471 speculator,
472 plan,
473 inputs,
474 config,
475 confirm_serial,
476 )
477 }
478
479 fn record_borrowed_batch_job(
480 &self,
481 program: &Program,
482 inputs: &[&[u8]],
483 config: &vyre_driver::DispatchConfig,
484 started: Instant,
485 ) -> Result<crate::engine::record_and_readback::RecordedDispatch, vyre_driver::BackendError>
486 {
487 self.enforce_config_caps(config)?;
488 self.validate_with_cache(program)?;
489 let pipeline = crate::pipeline::WgpuPipeline::compile_with_device_queue(
490 program,
491 config,
492 self.adapter_info.clone(),
493 self.enabled_features,
494 self.current_device_queue(),
495 self.dispatch_arena_snapshot(),
496 self.current_persistent_pool(),
497 self.pipeline_cache.clone(),
498 self.bind_group_layout_cache.clone(),
499 )?;
500 if let Some(deadline) = config.timeout {
501 let elapsed = started.elapsed();
502 if elapsed > deadline {
503 return Err(vyre_driver::BackendError::new(format!(
504 "batch dispatch cancelled before GPU submission: took {elapsed:?}, budget {deadline:?}. Fix: raise DispatchConfig.timeout or split the program into smaller chunks."
505 )));
506 }
507 }
508 let workgroup_count = pipeline.workgroups_for_dispatch(config)?;
509 let dispatch_arena = self.dispatch_arena_snapshot();
510 crate::engine::record_and_readback::record_dispatch_unsubmitted(
511 crate::engine::record_and_readback::RecordAndReadback::for_dispatch(
512 &pipeline,
513 &dispatch_arena,
514 inputs,
515 workgroup_count,
516 config,
517 crate::async_dispatch::timestamp_profile_requested(config),
518 crate::engine::record_and_readback::DispatchLabels {
519 bind_group: "vyre batch dispatch bind group",
520 encoder: "vyre batch dispatch",
521 compute: "vyre batch dispatch compute",
522 },
523 ),
524 )
525 }
526
527 pub fn dispatch_borrowed_batch(
529 &self,
530 jobs: &[(&Program, &[&[u8]], &vyre_driver::DispatchConfig)],
531 ) -> Result<
532 Vec<Result<vyre_driver::OutputBuffers, vyre_driver::BackendError>>,
533 vyre_driver::BackendError,
534 > {
535 let _span = tracing::trace_span!(
536 "vyre.dispatch_borrowed_batch",
537 backend = "wgpu",
538 jobs = jobs.len(),
539 );
540 let _enter = _span.enter();
541
542 let mut results = empty_batch_result_slots(jobs.len())?;
543 let mut recorded = Vec::new();
544 reserve_backend_vec(&mut recorded, jobs.len(), "recorded dispatch")?;
545 let mut meta = Vec::new();
546 reserve_backend_vec(&mut meta, jobs.len(), "batch dispatch metadata")?;
547 for (index, (program, inputs, config)) in jobs.iter().enumerate() {
548 let started = Instant::now();
549 if program.is_explicit_noop() {
550 results[index] = Some(Ok(Vec::new()));
551 continue;
552 }
553 let command = self.record_borrowed_batch_job(program, inputs, config, started)?;
554 recorded.push(command);
555 meta.push((index, started, config.timeout));
556 }
557
558 let pending = crate::engine::record_and_readback::submit_recorded_batch(recorded)?;
559 for ((index, started, timeout), result) in meta
560 .into_iter()
561 .zip(crate::engine::record_and_readback::WgpuPendingReadback::await_many_owned(pending))
562 {
563 results[index] = Some(result.and_then(|outputs| {
564 if let Some(deadline) = timeout {
565 let elapsed = started.elapsed();
566 if elapsed > deadline {
567 return Err(vyre_driver::BackendError::new(format!(
568 "batch dispatch exceeded configured timeout: took {elapsed:?}, budget {deadline:?}. Fix: raise DispatchConfig.timeout or split the program into smaller chunks."
569 )));
570 }
571 }
572 Ok(outputs)
573 }));
574 }
575 finalize_batch_results(
576 results,
577 "internal batch dispatch result slot was not filled. Fix: keep batch recording metadata synchronized.",
578 )
579 }
580
581 pub fn dispatch_borrowed_batch_into(
583 &self,
584 jobs: &[(&Program, &[&[u8]], &vyre_driver::DispatchConfig)],
585 outputs: &mut [vyre_driver::OutputBuffers],
586 ) -> Result<Vec<Result<(), vyre_driver::BackendError>>, vyre_driver::BackendError> {
587 if outputs.len() != jobs.len() {
588 return Err(vyre_driver::BackendError::new(format!(
589 "dispatch_borrowed_batch_into received {} output slots for {} jobs. Fix: pass exactly one OutputBuffers slot per job.",
590 outputs.len(),
591 jobs.len()
592 )));
593 }
594
595 let _span = tracing::trace_span!(
596 "vyre.dispatch_borrowed_batch_into",
597 backend = "wgpu",
598 jobs = jobs.len(),
599 );
600 let _enter = _span.enter();
601
602 let mut results = empty_batch_result_slots(jobs.len())?;
603 let mut recorded = Vec::new();
604 reserve_backend_vec(&mut recorded, jobs.len(), "recorded dispatch")?;
605 let mut meta = Vec::new();
606 reserve_backend_vec(&mut meta, jobs.len(), "batch-into dispatch metadata")?;
607 for (index, (program, inputs, config)) in jobs.iter().enumerate() {
608 let started = Instant::now();
609 if program.is_explicit_noop() {
610 outputs[index].clear();
611 results[index] = Some(Ok(()));
612 continue;
613 }
614 let command = self.record_borrowed_batch_job(program, inputs, config, started)?;
615 recorded.push(command);
616 meta.push((index, started, config.timeout));
617 }
618
619 let pending = crate::engine::record_and_readback::submit_recorded_batch(recorded)?;
620 let deadline =
621 crate::engine::record_and_readback::WgpuPendingReadback::wait_for_many(&pending);
622 for ((index, started, timeout), readback) in meta.into_iter().zip(pending) {
623 results[index] = Some(
624 readback
625 .collect_after_submission_wait(&mut outputs[index], deadline)
626 .and_then(|()| {
627 if let Some(deadline) = timeout {
628 let elapsed = started.elapsed();
629 if elapsed > deadline {
630 return Err(vyre_driver::BackendError::new(format!(
631 "batch dispatch exceeded configured timeout: took {elapsed:?}, budget {deadline:?}. Fix: raise DispatchConfig.timeout or split the program into smaller chunks."
632 )));
633 }
634 }
635 Ok(())
636 }),
637 );
638 }
639 finalize_batch_results(
640 results,
641 "internal batch-into dispatch result slot was not filled. Fix: keep batch recording metadata synchronized.",
642 )
643 }
644
645 pub fn dispatch_batch(
647 &self,
648 jobs: &[(
649 vyre_foundation::ir::Program,
650 Vec<Vec<u8>>,
651 vyre_driver::DispatchConfig,
652 )],
653 ) -> Result<
654 Vec<Result<vyre_driver::OutputBuffers, vyre_driver::BackendError>>,
655 vyre_driver::BackendError,
656 > {
657 let mut borrowed_inputs = Vec::new();
658 reserve_backend_vec(&mut borrowed_inputs, jobs.len(), "borrowed input batch")?;
659 for (_, inputs, _) in jobs {
660 let mut borrowed = smallvec::SmallVec::<[&[u8]; 8]>::new();
661 reserve_smallvec(
662 &mut borrowed,
663 inputs.len(),
664 "WGPU backend",
665 "borrowed input slice reference",
666 "split the batch job before dispatch",
667 )?;
668 borrowed.extend(inputs.iter().map(Vec::as_slice));
669 borrowed_inputs.push(borrowed);
670 }
671 let mut borrowed_jobs = Vec::new();
672 reserve_backend_vec(&mut borrowed_jobs, jobs.len(), "borrowed dispatch job")?;
673 for ((program, _, config), inputs) in jobs.iter().zip(borrowed_inputs.iter()) {
674 borrowed_jobs.push((program, inputs.as_slice(), config));
675 }
676 self.dispatch_borrowed_batch(&borrowed_jobs)
677 }
678
679 #[doc(hidden)]
681 pub fn compile_pipeline_for_oracle(
682 &self,
683 program: &vyre_foundation::ir::Program,
684 config: &vyre_driver::DispatchConfig,
685 ) -> Result<Arc<crate::pipeline::WgpuPipeline>, vyre_driver::BackendError> {
686 self.enforce_config_caps(config)?;
687 crate::pipeline::WgpuPipeline::compile_with_device_queue(
688 program,
689 config,
690 self.adapter_info.clone(),
691 self.enabled_features,
692 self.current_device_queue(),
693 self.dispatch_arena_snapshot(),
694 self.current_persistent_pool(),
695 self.pipeline_cache.clone(),
696 self.bind_group_layout_cache.clone(),
697 )
698 }
699}
700
701#[allow(clippy::needless_lifetimes)]
708pub(crate) fn borrowed_slices_from_owned_inputs<'a>(
709 inputs: &'a [Vec<u8>],
710) -> smallvec::SmallVec<[&'a [u8]; 8]> {
711 let mut borrowed = smallvec::SmallVec::<[&'a [u8]; 8]>::with_capacity(inputs.len());
712 borrowed.extend(inputs.iter().map(Vec::as_slice));
713 borrowed
714}
715
716impl vyre_driver::VyreBackend for WgpuBackend {
717 fn id(&self) -> &'static str {
718 "wgpu"
719 }
720
721 fn version(&self) -> &'static str {
722 env!("CARGO_PKG_VERSION")
723 }
724
725 fn supported_ops(&self) -> &std::collections::HashSet<vyre_foundation::ir::OpId> {
726 vyre_driver::backend::validation::default_supported_ops_with_trap()
727 }
728
729 fn dispatch(
730 &self,
731 program: &Program,
732 inputs: &[Vec<u8>],
733 config: &vyre_driver::DispatchConfig,
734 ) -> Result<Vec<Vec<u8>>, vyre_driver::BackendError> {
735 let _span = tracing::trace_span!(
736 "vyre.dispatch",
737 backend = "wgpu",
738 inputs = inputs.len(),
739 label = tracing::field::Empty,
740 );
741 let _enter = _span.enter();
742 if let Some(label) = config.label.as_deref() {
743 _span.record("label", label);
744 }
745 let borrowed = borrowed_slices_from_owned_inputs(inputs);
746 let start = Instant::now();
747 let result = self
748 .dispatch_borrowed_async(program, &borrowed, config)?
749 .await_owned();
750 tracing::trace!(
751 target: "vyre.dispatch",
752 elapsed_us = elapsed_micros_u64(start, "borrowed-path dispatch")?,
753 inputs = inputs.len(),
754 "dispatch completed (borrowed-path; clone-free)"
755 );
756 result
757 }
758
759 fn dispatch_borrowed(
760 &self,
761 program: &Program,
762 inputs: &[&[u8]],
763 config: &vyre_driver::DispatchConfig,
764 ) -> Result<Vec<Vec<u8>>, vyre_driver::BackendError> {
765 let _span = tracing::trace_span!(
766 "vyre.dispatch",
767 backend = "wgpu",
768 inputs = inputs.len(),
769 label = tracing::field::Empty,
770 );
771 let _enter = _span.enter();
772 if let Some(label) = config.label.as_deref() {
773 _span.record("label", label);
774 }
775 let start = Instant::now();
776 let result = self
777 .dispatch_borrowed_async(program, inputs, config)?
778 .await_owned();
779 tracing::trace!(
780 target: "vyre.dispatch",
781 elapsed_us = elapsed_micros_u64(start, "dispatch")?,
782 inputs = inputs.len(),
783 "dispatch completed"
784 );
785 result
786 }
787
788 fn dispatch_borrowed_into(
789 &self,
790 program: &Program,
791 inputs: &[&[u8]],
792 config: &vyre_driver::DispatchConfig,
793 outputs: &mut vyre_driver::OutputBuffers,
794 ) -> Result<(), vyre_driver::BackendError> {
795 let _span = tracing::trace_span!(
796 "vyre.dispatch_into",
797 backend = "wgpu",
798 inputs = inputs.len(),
799 label = tracing::field::Empty,
800 );
801 let _enter = _span.enter();
802 if let Some(label) = config.label.as_deref() {
803 _span.record("label", label);
804 }
805 if vyre_driver::grid_sync::contains_grid_sync(program)
806 && !<Self as vyre_driver::VyreBackend>::supports_grid_sync(self)
807 {
808 return vyre_driver::grid_sync::dispatch_with_grid_sync_split_into(
809 self, program, inputs, config, outputs,
810 );
811 }
812 let start = Instant::now();
813 self.dispatch_borrowed_async(program, inputs, config)?
814 .await_into(outputs)?;
815 tracing::trace!(
816 target: "vyre.dispatch",
817 elapsed_us = elapsed_micros_u64(start, "dispatch into caller-owned outputs")?,
818 inputs = inputs.len(),
819 "dispatch completed into caller-owned outputs"
820 );
821 Ok(())
822 }
823
824 fn dispatch_borrowed_timed(
825 &self,
826 program: &Program,
827 inputs: &[&[u8]],
828 config: &vyre_driver::DispatchConfig,
829 ) -> Result<vyre_driver::TimedDispatchResult, vyre_driver::BackendError> {
830 let _span = tracing::trace_span!(
831 "vyre.dispatch_timed",
832 backend = "wgpu",
833 inputs = inputs.len(),
834 label = tracing::field::Empty,
835 );
836 let _enter = _span.enter();
837 if let Some(label) = config.label.as_deref() {
838 _span.record("label", label);
839 }
840 WgpuBackend::dispatch_borrowed_async_timed(self, program, inputs, config)?
841 .await_timed_owned()
842 }
843
844 fn dispatch_async(
845 &self,
846 program: &Program,
847 inputs: &[Vec<u8>],
848 config: &vyre_driver::DispatchConfig,
849 ) -> Result<Box<dyn vyre_driver::backend::PendingDispatch>, vyre_driver::BackendError> {
850 let _span = tracing::trace_span!(
851 "vyre.dispatch_async",
852 backend = "wgpu",
853 inputs = inputs.len(),
854 label = tracing::field::Empty,
855 );
856 let _enter = _span.enter();
857 if let Some(label) = config.label.as_deref() {
858 _span.record("label", label);
859 }
860
861 let borrowed = borrowed_slices_from_owned_inputs(inputs);
862 Ok(Box::new(WgpuBackend::dispatch_borrowed_async(
863 self, program, &borrowed, config,
864 )?))
865 }
866
867 fn dispatch_borrowed_async(
868 &self,
869 program: &Program,
870 inputs: &[&[u8]],
871 config: &vyre_driver::DispatchConfig,
872 ) -> Result<Box<dyn vyre_driver::backend::PendingDispatch>, vyre_driver::BackendError> {
873 Ok(Box::new(WgpuBackend::dispatch_borrowed_async(
874 self, program, inputs, config,
875 )?))
876 }
877
878 fn allocate_device_buffer(
879 &self,
880 byte_len: usize,
881 ) -> Result<Box<dyn vyre_driver::DeviceBuffer>, vyre_driver::BackendError> {
882 self.allocate_wgpu_device_buffer(byte_len)
883 }
884
885 fn upload_device_buffer(
886 &self,
887 buffer: &mut dyn vyre_driver::DeviceBuffer,
888 bytes: &[u8],
889 ) -> Result<(), vyre_driver::BackendError> {
890 self.upload_wgpu_device_buffer(buffer, bytes)
891 }
892
893 fn download_device_buffer(
894 &self,
895 buffer: &dyn vyre_driver::DeviceBuffer,
896 ) -> Result<Vec<u8>, vyre_driver::BackendError> {
897 self.download_wgpu_device_buffer(buffer)
898 }
899
900 fn free_device_buffer(
901 &self,
902 buffer: Box<dyn vyre_driver::DeviceBuffer>,
903 ) -> Result<(), vyre_driver::BackendError> {
904 self.free_wgpu_device_buffer(buffer)
905 }
906
907 fn allocate_resident(
908 &self,
909 byte_len: usize,
910 ) -> Result<vyre_driver::Resource, vyre_driver::BackendError> {
911 crate::resident_resource::allocate_resident(self, byte_len)
912 }
913
914 fn upload_resident(
915 &self,
916 resource: &vyre_driver::Resource,
917 bytes: &[u8],
918 ) -> Result<(), vyre_driver::BackendError> {
919 crate::resident_upload::upload_resident(self, resource, bytes)
920 }
921
922 fn upload_resident_many(
923 &self,
924 uploads: &[(&vyre_driver::Resource, &[u8])],
925 ) -> Result<(), vyre_driver::BackendError> {
926 crate::resident_upload::upload_resident_many(self, uploads)
927 }
928
929 fn upload_resident_at(
930 &self,
931 resource: &vyre_driver::Resource,
932 dst_offset_bytes: usize,
933 bytes: &[u8],
934 ) -> Result<(), vyre_driver::BackendError> {
935 crate::resident_upload::upload_resident_at(self, resource, dst_offset_bytes, bytes)
936 }
937
938 fn upload_resident_at_many(
939 &self,
940 uploads: &[(&vyre_driver::Resource, usize, &[u8])],
941 ) -> Result<(), vyre_driver::BackendError> {
942 crate::resident_upload::upload_resident_at_many(self, uploads)
943 }
944
945 fn download_resident(
946 &self,
947 resource: &vyre_driver::Resource,
948 ) -> Result<Vec<u8>, vyre_driver::BackendError> {
949 crate::resident_download::download_resident(self, resource)
950 }
951
952 fn download_resident_into(
953 &self,
954 resource: &vyre_driver::Resource,
955 out: &mut Vec<u8>,
956 ) -> Result<(), vyre_driver::BackendError> {
957 crate::resident_download::download_resident_into(self, resource, out)
958 }
959
960 fn download_resident_range(
961 &self,
962 resource: &vyre_driver::Resource,
963 byte_offset: usize,
964 byte_len: usize,
965 ) -> Result<Vec<u8>, vyre_driver::BackendError> {
966 crate::resident_download::download_resident_range(self, resource, byte_offset, byte_len)
967 }
968
969 fn download_resident_range_into(
970 &self,
971 resource: &vyre_driver::Resource,
972 byte_offset: usize,
973 byte_len: usize,
974 out: &mut Vec<u8>,
975 ) -> Result<(), vyre_driver::BackendError> {
976 crate::resident_download::download_resident_range_into(
977 self,
978 resource,
979 byte_offset,
980 byte_len,
981 out,
982 )
983 }
984
985 fn download_resident_ranges_into(
986 &self,
987 ranges: &[(&vyre_driver::Resource, usize, usize)],
988 outputs: &mut [&mut Vec<u8>],
989 ) -> Result<(), vyre_driver::BackendError> {
990 crate::resident_download::download_resident_ranges_into(self, ranges, outputs)
991 }
992
993 fn free_resident(
994 &self,
995 resource: vyre_driver::Resource,
996 ) -> Result<(), vyre_driver::BackendError> {
997 crate::resident_resource::free_resident(self, resource)
998 }
999
1000 fn dispatch_resident_timed(
1001 &self,
1002 program: &Program,
1003 resources: &[vyre_driver::Resource],
1004 config: &vyre_driver::DispatchConfig,
1005 ) -> Result<vyre_driver::TimedDispatchResult, vyre_driver::BackendError> {
1006 crate::resident_dispatch::dispatch_resident_timed(self, program, resources, config)
1007 }
1008
1009 fn dispatch_resident_async(
1010 &self,
1011 program: &Program,
1012 resources: &[vyre_driver::Resource],
1013 config: &vyre_driver::DispatchConfig,
1014 ) -> Result<Box<dyn vyre_driver::PendingDispatch>, vyre_driver::BackendError> {
1015 crate::resident_dispatch::dispatch_resident_async(self, program, resources, config)
1016 }
1017
1018 fn dispatch_with_device_buffers(
1019 &self,
1020 program: &Program,
1021 inputs: &[&dyn vyre_driver::DeviceBuffer],
1022 outputs: &mut [&mut dyn vyre_driver::DeviceBuffer],
1023 config: &vyre_driver::DispatchConfig,
1024 ) -> Result<(), vyre_driver::BackendError> {
1025 vyre_driver::validate_buffer_ownership(self.id(), inputs.iter().copied())?;
1028 vyre_driver::validate_buffer_ownership(
1029 self.id(),
1030 outputs
1031 .iter()
1032 .map(|b| &**b as &dyn vyre_driver::DeviceBuffer),
1033 )?;
1034
1035 let resource_count = inputs.len().checked_add(outputs.len()).ok_or_else(|| {
1036 vyre_driver::BackendError::new(
1037 "resident dispatch resource count overflowed usize. Fix: split input/output resources before dispatch.",
1038 )
1039 })?;
1040 let mut resources =
1041 smallvec::SmallVec::<[vyre_driver::Resource; 8]>::with_capacity(resource_count);
1042 for buffer in inputs {
1043 let wgpu_buf = buffer
1044 .as_any()
1045 .downcast_ref::<crate::WgpuDeviceBuffer>()
1046 .ok_or_else(|| {
1047 vyre_driver::BackendError::new(format!(
1048 "Fix: dispatch_with_device_buffers expected WgpuDeviceBuffer inputs but got buffer owned by `{}`.",
1049 buffer.backend_id()
1050 ))
1051 })?;
1052 resources.push(vyre_driver::Resource::Resident(
1053 wgpu_buf.handle().resident_handle()?,
1054 ));
1055 }
1056 for buffer in outputs.iter() {
1057 let backend_id = buffer.backend_id().to_string();
1058 let wgpu_buf = buffer
1059 .as_any()
1060 .downcast_ref::<crate::WgpuDeviceBuffer>()
1061 .ok_or_else(|| {
1062 vyre_driver::BackendError::new(format!(
1063 "Fix: dispatch_with_device_buffers expected WgpuDeviceBuffer outputs but got buffer owned by `{backend_id}`."
1064 ))
1065 })?;
1066 resources.push(vyre_driver::Resource::Resident(
1067 wgpu_buf.handle().resident_handle()?,
1068 ));
1069 }
1070
1071 let pipeline = self.compile_resident_pipeline_cached(program, config)?;
1072 let _outputs = vyre_driver::CompiledPipeline::dispatch_persistent_handles(
1073 &*pipeline, &resources, config,
1074 )?;
1075 Ok(())
1076 }
1077
1078 fn pipeline_cache_snapshot(&self) -> Option<vyre_driver::pipeline::PipelineCacheSnapshot> {
1079 Some(vyre_driver::pipeline::PipelineCacheSnapshot {
1080 hits: self.pipeline_cache.hits(),
1081 misses: self.pipeline_cache.misses(),
1082 })
1083 }
1084
1085 fn supports_subgroup_ops(&self) -> bool {
1086 crate::capabilities::supports_subgroup_ops(&self.enabled_features)
1087 }
1088
1089 fn supports_f16(&self) -> bool {
1090 false
1091 }
1092
1093 fn supports_bf16(&self) -> bool {
1094 false
1095 }
1096
1097 fn supports_tensor_cores(&self) -> bool {
1098 false
1099 }
1100
1101 fn supports_async_compute(&self) -> bool {
1102 false
1103 }
1104
1105 fn supports_indirect_dispatch(&self) -> bool {
1106 crate::capabilities::supports_indirect_dispatch(&self.adapter_info, &self.enabled_features)
1107 }
1108
1109 fn supports_speculation(&self) -> bool {
1110 false
1111 }
1112
1113 fn supports_persistent_thread_dispatch(&self) -> bool {
1114 false
1115 }
1116
1117 fn is_distributed(&self) -> bool {
1118 false
1119 }
1120
1121 fn max_workgroup_size(&self) -> [u32; 3] {
1122 self.enabled_features.max_workgroup_size
1123 }
1124
1125 fn max_compute_workgroups_per_dimension(&self) -> u32 {
1126 self.device_limits.max_compute_workgroups_per_dimension
1127 }
1128
1129 fn max_compute_invocations_per_workgroup(&self) -> u32 {
1130 self.device_limits.max_compute_invocations_per_workgroup
1131 }
1132
1133 fn subgroup_size(&self) -> Option<u32> {
1134 crate::capabilities::supports_subgroup_ops(&self.enabled_features)
1135 .then_some(self.enabled_features.min_subgroup_size)
1136 }
1137
1138 fn max_storage_buffer_bytes(&self) -> u64 {
1139 self.enabled_features.max_storage_buffer_binding_size
1140 }
1141
1142 fn device_profile(&self) -> vyre_driver::DeviceProfile {
1143 WgpuBackend::device_profile(self)
1144 }
1145
1146 fn flush(&self) -> Result<(), vyre_driver::BackendError> {
1147 let device_queue = self.current_device_queue();
1148 let submission = device_queue.1.submit(std::iter::empty());
1149 crate::runtime::device::poll_device_wait_for(&device_queue.0, submission)?;
1150 crate::pipeline::disk_cache::flush_disk_pipeline_cache()
1151 }
1152
1153 fn device_lost(&self) -> bool {
1154 self.device_lost.load(Ordering::Acquire)
1155 }
1156
1157 fn try_recover(&self) -> Result<(), vyre_driver::BackendError> {
1158 let ((device, queue), adapter_info, enabled) = match &self.recovery_target {
1159 AdapterRecoveryTarget::Index(index) => {
1160 crate::runtime::device::init_device_for_adapter(*index)
1161 }
1162 AdapterRecoveryTarget::Identity(identity) => {
1163 crate::runtime::device::init_device_for_adapter_identity(identity)
1164 }
1165 }
1166 .map_err(|error| vyre_driver::BackendError::new(error.to_string()))?;
1167 let device_limits = device.limits();
1168 let recovered_identity = crate::runtime::device::AdapterIdentity::from_info(&adapter_info);
1169 let original_identity =
1170 crate::runtime::device::AdapterIdentity::from_info(&self.adapter_info);
1171 if recovered_identity != original_identity {
1172 return Err(vyre_driver::BackendError::new(format!(
1173 "wgpu recovery selected a different adapter than the backend was constructed with. Original: {:?}; recovered: {:?}. Fix: construct a new backend for the new adapter instead of reusing device-local caches across adapter identities.",
1174 self.adapter_info, adapter_info
1175 )));
1176 }
1177 if device_limits != self.device_limits || enabled != self.enabled_features {
1178 return Err(vyre_driver::BackendError::new(
1179 "wgpu recovery selected the original adapter but feature or limit negotiation changed. Fix: construct a new backend so dispatch planning and pipeline caches are rebuilt against the new device contract.",
1180 ));
1181 }
1182 let cache_tiers = vec![
1183 crate::runtime::cache::CacheTier::try_new("hot", 1 << 24)?,
1184 crate::runtime::cache::CacheTier::try_new("cold", 1 << 30)?,
1185 ];
1186 let persistent_pool = crate::buffer::BufferPool::with_tiering(
1187 device.clone(),
1188 queue.clone(),
1189 &vyre_driver::DispatchConfig::default(),
1190 cache_tiers,
1191 )?;
1192 self.device_queue
1193 .store(Arc::new((device.clone(), queue.clone())));
1194 self.persistent_pool.store(Arc::new(persistent_pool));
1195 self.pipeline_cache.clear();
1196 self.wgsl_dispatch_pipeline_cache.clear();
1197 self.bind_group_layout_cache.clear();
1198 self.validation_cache.clear()?;
1199 self.dispatch_arena.store(Arc::new(DispatchArena::new(
1200 device.clone(),
1201 queue.clone(),
1202 &vyre_driver::DispatchConfig::default(),
1203 )));
1204 self.device_lost.store(false, Ordering::Release);
1205
1206 Ok(())
1207 }
1208}
1209
1210impl vyre_self_substrate::optimizer::dispatcher::OptimizerDispatcher for WgpuBackend {
1211 fn dispatch(
1212 &self,
1213 program: &Program,
1214 inputs: &[Vec<u8>],
1215 grid_override: Option<[u32; 3]>,
1216 ) -> Result<Vec<Vec<u8>>, vyre_self_substrate::optimizer::dispatcher::DispatchError> {
1217 let mut config = vyre_driver::DispatchConfig::default();
1218 config.grid_override = grid_override;
1219 vyre_driver::VyreBackend::dispatch(self, program, inputs, &config).map_err(|error| {
1220 vyre_self_substrate::optimizer::dispatcher::DispatchError::BackendError(
1221 error.to_string(),
1222 )
1223 })
1224 }
1225}
1226
1227#[cfg(test)]
1228mod borrowed_slice_conversion_tests {
1229 use super::{
1230 borrowed_slices_from_owned_inputs, empty_batch_result_slots, finalize_batch_results,
1231 };
1232
1233 #[test]
1234 fn dispatch_async_input_conversion_is_zero_copy_slice_refs() {
1235 let inputs = vec![vec![1u8, 2, 3], vec![4u8, 5]];
1236 let borrowed = borrowed_slices_from_owned_inputs(&inputs);
1237 assert_eq!(borrowed.len(), 2);
1238 assert_eq!(borrowed[0].as_ptr(), inputs[0].as_ptr());
1239 assert_eq!(borrowed[1].as_ptr(), inputs[1].as_ptr());
1240 }
1241
1242 #[test]
1243 fn nine_inputs_spill_smallvec_but_slices_alias_vecs() {
1244 let inputs: Vec<Vec<u8>> = (0..9).map(|i| vec![i as u8]).collect();
1245 let borrowed = borrowed_slices_from_owned_inputs(&inputs);
1246 assert_eq!(borrowed.len(), 9);
1247 for i in 0..9 {
1248 assert_eq!(
1249 borrowed[i].as_ptr(),
1250 inputs[i].as_ptr(),
1251 "slice {i} must reference the corresponding Vec buffer"
1252 );
1253 }
1254 }
1255
1256 #[test]
1257 fn generated_batch_result_finalization_preserves_success_error_and_missing_slots() {
1258 for case in 0..4096usize {
1259 let len = (case % 19) + 1;
1260 let mut slots = empty_batch_result_slots::<usize>(len)
1261 .expect("Fix: generated WGPU batch result test must reserve slots");
1262 for slot in 0..len {
1263 match (slot + case) % 7 {
1264 0 => {}
1265 1 => {
1266 slots[slot] = Some(Err(vyre_driver::BackendError::new(format!(
1267 "generated-error-{case}-{slot}"
1268 ))));
1269 }
1270 _ => {
1271 slots[slot] = Some(Ok(case * 100 + slot));
1272 }
1273 }
1274 }
1275
1276 let finalized =
1277 finalize_batch_results(slots, "generated missing WGPU batch result slot")
1278 .expect("Fix: generated WGPU batch finalization must reserve output results");
1279 assert_eq!(
1280 finalized.len(),
1281 len,
1282 "generated WGPU batch case {case} must preserve slot count"
1283 );
1284 for (slot, result) in finalized.into_iter().enumerate() {
1285 match (slot + case) % 7 {
1286 0 => {
1287 let error =
1288 result.expect_err("Fix: missing generated batch slot must error");
1289 assert!(
1290 error
1291 .to_string()
1292 .contains("generated missing WGPU batch result slot"),
1293 "Fix: missing generated batch slot must report the supplied invariant, got {error}"
1294 );
1295 }
1296 1 => {
1297 let error = result
1298 .expect_err("Fix: explicit generated batch error must stay error");
1299 assert!(
1300 error
1301 .to_string()
1302 .contains(&format!("generated-error-{case}-{slot}")),
1303 "Fix: explicit generated batch error must be preserved, got {error}"
1304 );
1305 }
1306 _ => {
1307 assert_eq!(
1308 result.expect("Fix: generated batch success must stay success"),
1309 case * 100 + slot
1310 );
1311 }
1312 }
1313 }
1314 }
1315 }
1316}