1use laddu_compile::{CompiledModel, ReductionPlan};
2#[cfg(feature = "wgpu")]
3use laddu_data::BatchLayout;
4use laddu_data::data::Dataset;
5use laddu_data::data::EventBatch;
6#[cfg(feature = "wgpu")]
7use laddu_data::data::{CacheStorage, MemoryPolicy, accurate::AccurateF64};
8#[cfg(feature = "wgpu")]
9use laddu_data::schema::Precision as DataPrecision;
10use laddu_expr::parameters::ParamValues;
11#[cfg(feature = "wgpu")]
12use laddu_memory::{MemoryDecision, MemoryFootprint};
13use num::complex::Complex64;
14use std::sync::Arc;
15
16#[cfg(feature = "wgpu")]
17use crate::preparation::{DatasetPreparation, DatasetStatsAccumulator, RuntimePreparationPlan};
18use crate::{
19 CpuBackend, CpuPlan, CpuPreparedDataset, Execution, PreparedDatasetStats, ReductionEvaluation,
20 RuntimeError, RuntimeResult,
21};
22
23#[derive(Clone, Debug)]
25pub enum PreparedModel {
26 Cpu(Arc<CpuPlan>),
28 #[cfg(feature = "wgpu")]
29 Wgpu(WgpuPlan),
31}
32
33#[derive(Clone, Debug)]
35pub enum PreparedDataset {
36 Cpu(CpuPreparedDataset),
38 #[cfg(feature = "wgpu")]
39 Wgpu(WgpuPreparedDataset),
41}
42
43impl PreparedDataset {
44 pub fn stats(&self) -> &PreparedDatasetStats {
46 match self {
47 Self::Cpu(dataset) => dataset.stats(),
48 #[cfg(feature = "wgpu")]
49 Self::Wgpu(dataset) => dataset.stats(),
50 }
51 }
52}
53
54impl PreparedModel {
55 pub fn required_event_scalars(&self) -> &[String] {
59 match self {
60 Self::Cpu(plan) => plan.required_event_scalars(),
61 #[cfg(feature = "wgpu")]
62 Self::Wgpu(plan) => &plan.required_event_scalars,
63 }
64 }
65
66 #[doc(hidden)]
68 pub fn batch_memory_estimate(&self, events: usize) -> usize {
69 let cache = match self {
70 Self::Cpu(plan) => plan.cache_memory_estimate(events),
71 #[cfg(feature = "wgpu")]
72 Self::Wgpu(_) => 0,
73 };
74 cache.saturating_add(events.saturating_mul(64))
76 }
77
78 #[doc(hidden)]
84 pub fn visit_batch_many<F>(
85 &self,
86 execution: &Execution,
87 parameters: &[ParamValues],
88 batch: &EventBatch,
89 mut consume: F,
90 ) -> RuntimeResult<()>
91 where
92 F: FnMut(usize, &[Complex64]) -> RuntimeResult<()> + Send,
93 {
94 match self {
95 Self::Cpu(plan) => {
96 let cached = crate::CpuCachedBatch::from_cache(plan.cache_event_batch(batch)?);
97 execution.install(|| {
98 for (index, parameters) in parameters.iter().enumerate() {
99 let values =
100 plan.evaluate_prepared_batch(execution, parameters, &cached, true)?;
101 consume(index, &values)?;
102 }
103 Ok(())
104 })
105 }
106 #[cfg(feature = "wgpu")]
107 Self::Wgpu(_) => {
108 for (index, parameters) in parameters.iter().enumerate() {
109 let values = self.evaluate_batch(parameters, batch)?;
110 consume(index, &values)?;
111 }
112 Ok(())
113 }
114 }
115 }
116
117 pub fn evaluate_batch(
124 &self,
125 params: &ParamValues,
126 batch: &EventBatch,
127 ) -> RuntimeResult<Vec<Complex64>> {
128 match self {
129 Self::Cpu(plan) => plan.evaluate_batch(params, batch),
130 #[cfg(feature = "wgpu")]
131 Self::Wgpu(plan) => plan
132 .kernel
133 .evaluate_batch(&plan.context, params, batch)
134 .map(|values| {
135 values
136 .into_iter()
137 .map(|(re, im)| Complex64::new(re, im))
138 .collect()
139 })
140 .map_err(wgpu_error),
141 }
142 }
143
144 pub(crate) fn evaluate_batch_outputs(
148 &self,
149 params: &ParamValues,
150 batch: &EventBatch,
151 outputs: &[laddu_expr::ExprId],
152 ) -> RuntimeResult<Vec<Vec<Complex64>>> {
153 match self {
154 Self::Cpu(plan) => plan.evaluate_batch_outputs(params, batch, outputs),
155 #[cfg(feature = "wgpu")]
156 Self::Wgpu(_) => Err(RuntimeError::Wgpu(
157 "multi-output query evaluation is not implemented by the WGPU backend".into(),
158 )),
159 }
160 }
161
162 pub fn evaluate_batch_with_gradient(
169 &self,
170 params: &ParamValues,
171 batch: &EventBatch,
172 ) -> RuntimeResult<Vec<crate::ValueGradient>> {
173 match self {
174 Self::Cpu(plan) => plan.evaluate_batch_with_gradient(params, batch),
175 #[cfg(feature = "wgpu")]
176 Self::Wgpu(_) => Err(RuntimeError::Wgpu(
177 "event-wise model gradients are not implemented by the WGPU backend".into(),
178 )),
179 }
180 }
181
182 pub fn prepare(model: &CompiledModel, execution: &Execution) -> RuntimeResult<Self> {
189 #[cfg(feature = "wgpu")]
190 if let Some(context) = execution.wgpu_context() {
191 return Ok(Self::Wgpu(WgpuPlan {
192 context: context.clone(),
193 preparation_params: model.params().default_values(),
194 kernel: std::sync::Arc::new(
195 laddu_wgpu::WgpuScalarKernel::compile(context, model).map_err(wgpu_error)?,
196 ),
197 required_event_scalars: crate::required_event_scalars(model),
198 }));
199 }
200 Ok(Self::Cpu(
201 CpuBackend.prepare_shared_for_execution(model, execution)?,
202 ))
203 }
204
205 pub fn prepare_dataset(
212 &self,
213 execution: &Execution,
214 dataset: &Dataset,
215 ) -> RuntimeResult<PreparedDataset> {
216 match self {
217 Self::Cpu(plan) => Ok(PreparedDataset::Cpu(
218 plan.prepare_dataset(execution, dataset)?,
219 )),
220 #[cfg(feature = "wgpu")]
221 Self::Wgpu(plan) => plan
222 .prepare_dataset(execution, dataset)
223 .map(PreparedDataset::Wgpu),
224 }
225 }
226
227 pub fn evaluate_prepared(
238 &self,
239 execution: &Execution,
240 params: &ParamValues,
241 dataset: &PreparedDataset,
242 source: &Dataset,
243 ) -> RuntimeResult<Vec<Complex64>> {
244 let mut output =
245 self.evaluate_prepared_many(execution, std::slice::from_ref(params), dataset, source)?;
246 Ok(output.pop().unwrap_or_default())
247 }
248
249 pub fn evaluate_prepared_many(
258 &self,
259 execution: &Execution,
260 params: &[ParamValues],
261 dataset: &PreparedDataset,
262 source: &Dataset,
263 ) -> RuntimeResult<Vec<Vec<Complex64>>> {
264 #[cfg(not(feature = "wgpu"))]
265 let _ = source;
266 #[allow(unreachable_patterns)]
267 match (self, dataset) {
268 (Self::Cpu(plan), PreparedDataset::Cpu(dataset)) => {
269 plan.evaluate_prepared_dataset_many(execution, params, dataset)
270 }
271 #[cfg(feature = "wgpu")]
272 (Self::Wgpu(_), PreparedDataset::Wgpu(_)) => {
273 evaluate_source_many(self, execution, params, source, None)
274 .map(|(values, _)| values)
275 }
276 _ => Err(RuntimeError::InvalidShape {
277 index: 0,
278 message: "prepared model and dataset use different backends".into(),
279 }),
280 }
281 }
282
283 #[doc(hidden)]
286 pub fn visit_prepared_many<F>(
287 &self,
288 execution: &Execution,
289 parameter_sets: &[(&ParamValues, &str)],
290 dataset: &PreparedDataset,
291 source: &Dataset,
292 reduction: Option<ReductionPlan>,
293 consume: F,
294 ) -> RuntimeResult<Vec<f64>>
295 where
296 F: FnMut(usize, usize, &[Complex64]) -> RuntimeResult<()>,
297 {
298 #[cfg(not(feature = "wgpu"))]
299 let _ = source;
300 #[allow(unreachable_patterns)]
301 match (self, dataset) {
302 (Self::Cpu(plan), PreparedDataset::Cpu(dataset)) => plan.visit_prepared_dataset_many(
303 execution,
304 parameter_sets,
305 dataset,
306 reduction,
307 false,
308 consume,
309 ),
310 #[cfg(feature = "wgpu")]
311 (Self::Wgpu(_), PreparedDataset::Wgpu(_)) => {
312 visit_source_many(self, execution, parameter_sets, source, reduction, consume)
313 }
314 _ => Err(RuntimeError::InvalidShape {
315 index: 0,
316 message: "prepared model and dataset use different backends".into(),
317 }),
318 }
319 }
320
321 #[doc(hidden)]
323 pub fn visit_prepared_many_parallel<F>(
324 &self,
325 execution: &Execution,
326 parameter_sets: &[(&ParamValues, &str)],
327 dataset: &PreparedDataset,
328 source: &Dataset,
329 reduction: Option<ReductionPlan>,
330 consume: F,
331 ) -> RuntimeResult<Vec<f64>>
332 where
333 F: FnMut(usize, usize, &[Complex64]) -> RuntimeResult<()> + Send,
334 {
335 #[cfg(not(feature = "wgpu"))]
336 let _ = source;
337 #[allow(unreachable_patterns)]
338 match (self, dataset) {
339 (Self::Cpu(plan), PreparedDataset::Cpu(dataset)) => execution.install(|| {
340 plan.visit_prepared_dataset_many(
341 execution,
342 parameter_sets,
343 dataset,
344 reduction,
345 true,
346 consume,
347 )
348 }),
349 #[cfg(feature = "wgpu")]
350 (Self::Wgpu(_), PreparedDataset::Wgpu(_)) => {
351 visit_source_many(self, execution, parameter_sets, source, reduction, consume)
352 }
353 _ => Err(RuntimeError::InvalidShape {
354 index: 0,
355 message: "prepared model and dataset use different backends".into(),
356 }),
357 }
358 }
359
360 pub fn evaluate_prepared_many_with_reduction(
367 &self,
368 execution: &Execution,
369 params: &[ParamValues],
370 dataset: &PreparedDataset,
371 source: &Dataset,
372 reduction: ReductionPlan,
373 ) -> RuntimeResult<(Vec<Vec<Complex64>>, Vec<f64>)> {
374 #[cfg(not(feature = "wgpu"))]
375 let _ = source;
376 #[allow(unreachable_patterns)]
377 match (self, dataset) {
378 (Self::Cpu(plan), PreparedDataset::Cpu(dataset)) => plan
379 .evaluate_prepared_dataset_many_with_reduction(
380 execution, params, dataset, reduction,
381 ),
382 #[cfg(feature = "wgpu")]
383 (Self::Wgpu(_), PreparedDataset::Wgpu(_)) => {
384 let (values, sums) =
385 evaluate_source_many(self, execution, params, source, Some(reduction))?;
386 Ok((values, sums.unwrap_or_default()))
387 }
388 _ => Err(RuntimeError::InvalidShape {
389 index: 0,
390 message: "prepared model and dataset use different backends".into(),
391 }),
392 }
393 }
394
395 pub fn reduce(
402 &self,
403 execution: &Execution,
404 params: &ParamValues,
405 dataset: &PreparedDataset,
406 reduction: ReductionPlan,
407 ) -> RuntimeResult<f64> {
408 #[allow(unreachable_patterns)]
409 match (self, dataset) {
410 (Self::Cpu(plan), PreparedDataset::Cpu(dataset)) => {
411 plan.reduce(execution, params, dataset, reduction)
412 }
413 #[cfg(feature = "wgpu")]
414 (Self::Wgpu(plan), PreparedDataset::Wgpu(dataset)) => {
415 plan.reduce(execution, params, dataset, reduction)
416 }
417 _ => Err(RuntimeError::InvalidShape {
418 index: 0,
419 message: "prepared model and dataset use different backends".into(),
420 }),
421 }
422 }
423
424 pub fn reduce_with_gradient(
431 &self,
432 execution: &Execution,
433 params: &ParamValues,
434 dataset: &PreparedDataset,
435 reduction: ReductionPlan,
436 ) -> RuntimeResult<ReductionEvaluation> {
437 #[allow(unreachable_patterns)]
438 match (self, dataset) {
439 (Self::Cpu(plan), PreparedDataset::Cpu(dataset)) => {
440 plan.reduce_with_gradient(execution, params, dataset, reduction)
441 }
442 #[cfg(feature = "wgpu")]
443 (Self::Wgpu(plan), PreparedDataset::Wgpu(dataset)) => {
444 plan.reduce_with_gradient(execution, params, dataset, reduction)
445 }
446 _ => Err(RuntimeError::InvalidShape {
447 index: 0,
448 message: "prepared model and dataset use different backends".into(),
449 }),
450 }
451 }
452}
453
454#[cfg(feature = "wgpu")]
455type PreparedValuesWithSums = (Vec<Vec<Complex64>>, Option<Vec<f64>>);
456
457#[cfg(feature = "wgpu")]
458fn evaluate_source_many(
459 model: &PreparedModel,
460 execution: &Execution,
461 params: &[ParamValues],
462 source: &Dataset,
463 reduction: Option<ReductionPlan>,
464) -> RuntimeResult<PreparedValuesWithSums> {
465 let local = (|| {
466 let mut output = params.iter().map(|_| Vec::new()).collect::<Vec<_>>();
467 let mut sums = reduction.map(|_| {
468 params
469 .iter()
470 .map(|_| AccurateF64::zero())
471 .collect::<Vec<_>>()
472 });
473 for batch in source
474 .batches()
475 .map_err(|error| RuntimeError::Data(error.to_string()))?
476 {
477 let batch = batch.map_err(|error| RuntimeError::Data(error.to_string()))?;
478 for (index, (parameters, values)) in params.iter().zip(&mut output).enumerate() {
479 let batch_values = model.evaluate_batch(parameters, &batch)?;
480 if let (Some(reduction), Some(sums)) = (reduction, sums.as_mut()) {
481 for (row, value) in batch_values.iter().enumerate() {
482 sums[index].push(batch.weights_at(row) * reduction.apply(*value)?.value());
483 }
484 }
485 values.extend(batch_values);
486 }
487 }
488 Ok((
489 output,
490 sums.map(|sums| sums.into_iter().map(AccurateF64::finish).collect()),
491 ))
492 })();
493 if !execution.all_succeeded(local.is_ok()) {
494 return local.and(Err(RuntimeError::DistributedPeerFailure));
495 }
496 let (output, sums) = local?;
497 Ok((
498 output,
499 sums.map(|sums: Vec<f64>| sums.into_iter().map(|sum| execution.sum_f64(sum)).collect()),
500 ))
501}
502
503#[cfg(feature = "wgpu")]
504fn visit_source_many<F>(
505 model: &PreparedModel,
506 execution: &Execution,
507 parameter_sets: &[(&ParamValues, &str)],
508 source: &Dataset,
509 reduction: Option<ReductionPlan>,
510 mut consume: F,
511) -> RuntimeResult<Vec<f64>>
512where
513 F: FnMut(usize, usize, &[Complex64]) -> RuntimeResult<()>,
514{
515 let local = (|| {
516 let mut offset = 0;
517 let mut sums = reduction.map(|_| {
518 parameter_sets
519 .iter()
520 .map(|_| AccurateF64::zero())
521 .collect::<Vec<_>>()
522 });
523 for batch in source
524 .batches()
525 .map_err(|error| RuntimeError::Data(error.to_string()))?
526 {
527 let batch = batch.map_err(|error| RuntimeError::Data(error.to_string()))?;
528 for (index, &(parameters, context)) in parameter_sets.iter().enumerate() {
529 let values = model.evaluate_batch(parameters, &batch).map_err(|error| {
530 RuntimeError::Parameter(format!("{context} evaluation failed: {error}"))
531 })?;
532 if let (Some(reduction), Some(sums)) = (reduction, sums.as_mut()) {
533 for (row, value) in values.iter().enumerate() {
534 let reduced = reduction.apply(*value).map_err(|error| {
535 RuntimeError::Parameter(format!("{context} reduction failed: {error}"))
536 })?;
537 sums[index].push(batch.weights_at(row) * reduced.value());
538 }
539 }
540 consume(offset, index, &values)?;
541 }
542 offset += batch.len();
543 }
544 Ok(sums
545 .unwrap_or_default()
546 .into_iter()
547 .map(AccurateF64::finish)
548 .collect::<Vec<_>>())
549 })();
550 if !execution.all_succeeded(local.is_ok()) {
551 return local.and(Err(RuntimeError::DistributedPeerFailure));
552 }
553 Ok(local?
554 .into_iter()
555 .map(|sum| execution.sum_f64(sum))
556 .collect())
557}
558
559#[cfg(feature = "wgpu")]
560#[derive(Clone)]
562pub struct WgpuPlan {
563 context: std::sync::Arc<laddu_wgpu::WgpuContext>,
564 preparation_params: ParamValues,
565 kernel: std::sync::Arc<laddu_wgpu::WgpuScalarKernel>,
566 required_event_scalars: Vec<String>,
567}
568
569#[cfg(feature = "wgpu")]
570#[derive(Clone, Debug)]
571struct WgpuDatasetPlan {
572 read_plan: laddu_data::io::ReadPlan,
573 preparation_plan: RuntimePreparationPlan,
574 device_decision: MemoryDecision,
575}
576
577#[cfg(feature = "wgpu")]
578impl WgpuDatasetPlan {
579 fn resolve(
580 read_plan: laddu_data::io::ReadPlan,
581 memory_policy: MemoryPolicy,
582 local_event_limit: usize,
583 host_footprint: MemoryFootprint,
584 prepared_footprint: MemoryFootprint,
585 host_available: u64,
586 device_available: Option<u64>,
587 ) -> RuntimeResult<Self> {
588 let mut preparation_plan = RuntimePreparationPlan::new(read_plan, local_event_limit);
589 let host_decision = preparation_plan.fit_staging(
590 "WGPU host staging",
591 host_footprint,
592 host_available,
593 "bounded host staging",
594 )?;
595 let host_chunks = local_event_limit
596 .saturating_add(host_decision.chunk_events.saturating_sub(1))
597 / host_decision.chunk_events.max(1);
598 let resident_footprint = MemoryFootprint::fixed(prepared_footprint.fixed_bytes)
599 .checked_scale_usize(host_chunks)
600 .and_then(|fixed| {
601 fixed.checked_add(MemoryFootprint::per_event(
602 prepared_footprint.bytes_per_event,
603 ))
604 })
605 .map_err(|error| RuntimeError::Data(format!("GPU working-set overflow: {error}")))?;
606 let resident_peak = resident_footprint.peak_bytes(local_event_limit);
607 let device_available = device_available
608 .ok_or_else(|| RuntimeError::Wgpu("GPU execution has no device memory pool".into()))?;
609 preparation_plan.select_storage(
610 memory_policy,
611 "device",
612 resident_peak <= device_available,
613 resident_peak,
614 device_available,
615 )?;
616 let device_decision = if preparation_plan.storage() == CacheStorage::Resident {
617 preparation_plan.fit_resident(
618 "WGPU prepared dataset",
619 resident_footprint,
620 device_available,
621 "resident",
622 host_decision.chunk_events,
623 )?
624 } else {
625 preparation_plan.fit_staging(
626 "WGPU prepared dataset",
627 prepared_footprint,
628 device_available,
629 "streaming",
630 )?
631 };
632 let chunk_events = device_decision
633 .chunk_events
634 .min(host_decision.chunk_events)
635 .max(1);
636 preparation_plan.clamp_read_plan(read_plan.chunk_size, chunk_events);
637 Ok(Self {
638 read_plan: preparation_plan.read_plan(),
639 preparation_plan,
640 device_decision,
641 })
642 }
643
644 fn reserve_storage(
645 &mut self,
646 pool: Option<&crate::MemoryPool>,
647 ) -> RuntimeResult<Option<crate::MemoryLease>> {
648 self.preparation_plan.reserve_storage(pool, || {
649 RuntimeError::Wgpu("GPU execution has no device memory pool".into())
650 })?;
651 Ok(self.preparation_plan.take_memory_lease())
652 }
653}
654
655#[cfg(feature = "wgpu")]
656impl std::fmt::Debug for WgpuPlan {
657 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
658 formatter
659 .debug_struct("WgpuPlan")
660 .field("adapter", &self.context.info().name)
661 .finish_non_exhaustive()
662 }
663}
664
665#[cfg(feature = "wgpu")]
666#[derive(Clone)]
668pub enum WgpuPreparedDataset {
669 Resident {
671 batches: std::sync::Arc<[laddu_wgpu::WgpuPreparedBatch]>,
673 stats: PreparedDatasetStats,
675 memory_lease: crate::MemoryLease,
677 },
678 Streaming {
680 dataset: Dataset,
682 read_plan: laddu_data::io::ReadPlan,
684 workspace: std::sync::Arc<std::sync::Mutex<Option<laddu_wgpu::WgpuPreparedBatch>>>,
686 stats: PreparedDatasetStats,
688 transient_bytes: u64,
690 },
691}
692
693#[cfg(feature = "wgpu")]
694impl std::fmt::Debug for WgpuPreparedDataset {
695 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
696 formatter
697 .debug_struct("WgpuPreparedDataset")
698 .field("stats", self.stats())
699 .finish_non_exhaustive()
700 }
701}
702
703#[cfg(feature = "wgpu")]
704impl WgpuPreparedDataset {
705 pub fn stats(&self) -> &PreparedDatasetStats {
707 match self {
708 Self::Resident { stats, .. } | Self::Streaming { stats, .. } => stats,
709 }
710 }
711
712 fn try_for_each_prepared_batch<F>(
713 &self,
714 execution: &Execution,
715 context: &laddu_wgpu::WgpuContext,
716 kernel: &laddu_wgpu::WgpuScalarKernel,
717 preparation_params: &ParamValues,
718 mut consume: F,
719 ) -> RuntimeResult<()>
720 where
721 F: FnMut(&laddu_wgpu::WgpuPreparedBatch) -> RuntimeResult<()>,
722 {
723 match self {
724 Self::Resident { batches, .. } => {
725 for batch in batches.iter() {
726 consume(batch)?;
727 }
728 }
729 Self::Streaming {
730 dataset,
731 read_plan,
732 workspace,
733 transient_bytes,
734 ..
735 } => {
736 let _memory = execution
737 .device_memory()
738 .ok_or_else(|| {
739 RuntimeError::Wgpu("GPU execution has no device memory pool".into())
740 })?
741 .reserve(*transient_bytes)?;
742 let mut workspace = workspace.lock().map_err(|_| {
743 RuntimeError::Wgpu("streaming workspace lock is poisoned".into())
744 })?;
745 for batch in dataset
746 .stream_with_plan(*read_plan)
747 .map_err(|error| RuntimeError::Data(error.to_string()))?
748 {
749 let batch = batch.map_err(|error| RuntimeError::Data(error.to_string()))?;
750 if let Some(prepared) = workspace.as_mut() {
751 if !kernel
752 .refresh_batch(context, preparation_params, &batch, prepared)
753 .map_err(wgpu_error)?
754 {
755 *prepared = kernel
756 .prepare_batch(context, preparation_params, &batch)
757 .map_err(wgpu_error)?;
758 }
759 } else {
760 *workspace = Some(
761 kernel
762 .prepare_batch(context, preparation_params, &batch)
763 .map_err(wgpu_error)?,
764 );
765 }
766 consume(
767 workspace
768 .as_ref()
769 .expect("streaming workspace was initialized"),
770 )?;
771 }
772 }
773 }
774 Ok(())
775 }
776}
777
778#[cfg(feature = "wgpu")]
779impl WgpuPlan {
780 fn prepare_dataset(
781 &self,
782 execution: &Execution,
783 dataset: &Dataset,
784 ) -> RuntimeResult<WgpuPreparedDataset> {
785 let preparation = DatasetPreparation::new(execution, dataset);
786 let preparation_plan = preparation.runtime_plan()?;
787 let read_plan = preparation_plan.read_plan();
788 let local_event_limit = preparation_plan.event_limit();
789 let prepared_footprint = self
790 .kernel
791 .prepared_memory_footprint(&self.preparation_params)
792 .map_err(|error| RuntimeError::Data(format!("GPU working-set overflow: {error}")))?;
793 let schema = dataset
794 .schema()
795 .map_err(|error| RuntimeError::Data(error.to_string()))?;
796 let host_footprint = BatchLayout::from_schema(&schema)
797 .schema_working_set(DataPrecision::F64, 2)
798 .map_err(|error| RuntimeError::Data(format!("host working-set overflow: {error}")))?;
799 let mut plan = WgpuDatasetPlan::resolve(
800 read_plan,
801 dataset.memory_policy(),
802 local_event_limit,
803 host_footprint,
804 prepared_footprint,
805 execution.host_memory().remaining(),
806 execution.device_memory().map(|pool| pool.remaining()),
807 )?;
808 let memory_lease = plan.reserve_storage(execution.device_memory())?;
809 for decision in plan.preparation_plan.take_decisions().into_iter().rev() {
810 execution.record_memory_decision(decision);
811 }
812 let local = (|| {
813 let mut batches = Vec::new();
814 let mut stats = DatasetStatsAccumulator::new();
815 for batch in dataset
816 .stream_with_plan(plan.read_plan)
817 .map_err(|error| RuntimeError::Data(error.to_string()))?
818 {
819 let batch = batch.map_err(|error| RuntimeError::Data(error.to_string()))?;
820 stats.observe(&batch);
821 if plan.preparation_plan.storage() == CacheStorage::Resident {
822 batches.push(
823 self.kernel
824 .prepare_batch(&self.context, &self.preparation_params, &batch)
825 .map_err(wgpu_error)?,
826 );
827 }
828 }
829 Ok::<_, RuntimeError>((batches, stats.finish()))
830 })();
831 let (batches, local_stats) = preparation.coordinate(local)?;
832 let resident_bytes = batches
833 .iter()
834 .map(laddu_wgpu::WgpuPreparedBatch::resident_bytes)
835 .sum();
836 let stats =
837 preparation.finish_stats(local_stats, resident_bytes, plan.preparation_plan.storage());
838 Ok(match plan.preparation_plan.storage() {
839 CacheStorage::Resident => WgpuPreparedDataset::Resident {
840 batches: batches.into(),
841 stats,
842 memory_lease: memory_lease.ok_or_else(|| {
843 RuntimeError::Wgpu("resident GPU dataset did not reserve device memory".into())
844 })?,
845 },
846 CacheStorage::Streaming => WgpuPreparedDataset::Streaming {
847 dataset: dataset.clone(),
848 read_plan: plan.read_plan,
849 workspace: Default::default(),
850 stats,
851 transient_bytes: plan.device_decision.estimated_peak_bytes,
852 },
853 })
854 }
855
856 fn reduce(
857 &self,
858 execution: &Execution,
859 params: &ParamValues,
860 dataset: &WgpuPreparedDataset,
861 reduction: ReductionPlan,
862 ) -> RuntimeResult<f64> {
863 let mut total = AccurateF64::zero();
864 dataset.try_for_each_prepared_batch(
865 execution,
866 &self.context,
867 &self.kernel,
868 &self.preparation_params,
869 |batch| {
870 total.push(
871 self.kernel
872 .reduce_prepared_batch(&self.context, params, batch, reduction)
873 .map_err(wgpu_error)?,
874 );
875 Ok(())
876 },
877 )?;
878 Ok(execution.sum_f64(total.finish()))
879 }
880
881 fn reduce_with_gradient(
882 &self,
883 execution: &Execution,
884 params: &ParamValues,
885 dataset: &WgpuPreparedDataset,
886 reduction: ReductionPlan,
887 ) -> RuntimeResult<ReductionEvaluation> {
888 let mut total = AccurateF64::zero();
889 let mut gradient = (0..params.layout().n_free())
890 .map(|_| AccurateF64::zero())
891 .collect::<Vec<_>>();
892 let mut consume = |batch: &laddu_wgpu::WgpuPreparedBatch| -> RuntimeResult<()> {
893 let (value, values) = self
894 .kernel
895 .reduce_prepared_batch_with_gradient(&self.context, params, batch, reduction)
896 .map_err(wgpu_error)?;
897 total.push(value);
898 for (sum, value) in gradient.iter_mut().zip(values) {
899 sum.push(value);
900 }
901 Ok(())
902 };
903 dataset.try_for_each_prepared_batch(
904 execution,
905 &self.context,
906 &self.kernel,
907 &self.preparation_params,
908 &mut consume,
909 )?;
910 let gradient = gradient
911 .into_iter()
912 .map(|sum| execution.sum_f64(sum.finish()))
913 .collect();
914 Ok(ReductionEvaluation::new(
915 execution.sum_f64(total.finish()),
916 gradient,
917 ))
918 }
919}
920
921#[cfg(feature = "wgpu")]
922fn wgpu_error(error: laddu_wgpu::WgpuError) -> RuntimeError {
923 RuntimeError::Wgpu(error.to_string())
924}
925
926#[cfg(test)]
927mod prepared_evaluation_tests {
928 use std::{cell::RefCell, rc::Rc, sync::Arc};
929
930 use laddu_compile::{CompiledModel, ReductionPlan};
931 use laddu_data::{
932 data::{Dataset, EventBatch, OwnedEvent},
933 schema::Schema,
934 };
935 use laddu_expr::{event_scalar, parameter};
936 use num::complex::Complex64;
937
938 use super::PreparedModel;
939 use crate::{CpuOptions, Device, Execution, ExecutionOptions, JitPolicy, ThreadPolicy};
940
941 fn dataset(streaming: bool) -> Dataset {
942 let schema = Arc::new(
943 Schema::new(std::iter::empty::<&str>(), ["x"], false)
944 .expect("test schema should be valid"),
945 );
946 let first_batch = EventBatch::from_events(
947 schema.clone(),
948 [
949 OwnedEvent::new(vec![], vec![1.0]),
950 OwnedEvent::new(vec![], vec![2.0]),
951 ],
952 )
953 .expect("first test batch should be valid");
954 let second_batch = EventBatch::from_events(schema, [OwnedEvent::new(vec![], vec![3.0])])
955 .expect("second test batch should be valid");
956 let dataset = Dataset::from_batches(vec![first_batch, second_batch])
957 .expect("test dataset should be valid");
958 if streaming {
959 dataset.streaming()
960 } else {
961 dataset.fastest()
962 }
963 }
964
965 #[test]
966 fn prepared_evaluation_preserves_event_order_for_resident_and_streaming_data() {
967 let model =
968 CompiledModel::from_expr(&(event_scalar("x") + parameter!("offset", initial: 0.5)))
969 .expect("test model should compile");
970 let execution = Execution::default();
971 let prepared_model =
972 PreparedModel::prepare(&model, &execution).expect("model should prepare");
973 let parameters = [
974 model.params().values(&[0.5]).expect("first parameters"),
975 model.params().values(&[1.5]).expect("second parameters"),
976 ];
977 let expected = [
978 vec![
979 Complex64::new(1.5, 0.0),
980 Complex64::new(2.5, 0.0),
981 Complex64::new(3.5, 0.0),
982 ],
983 vec![
984 Complex64::new(2.5, 0.0),
985 Complex64::new(3.5, 0.0),
986 Complex64::new(4.5, 0.0),
987 ],
988 ];
989
990 for streaming in [false, true] {
991 let source = dataset(streaming);
992 let prepared = prepared_model
993 .prepare_dataset(&execution, &source)
994 .expect("dataset should prepare");
995 let actual = prepared_model
996 .evaluate_prepared_many(&execution, ¶meters, &prepared, &source)
997 .expect("prepared dataset should evaluate");
998 assert_eq!(actual, expected);
999 let (actual, sums) = prepared_model
1000 .evaluate_prepared_many_with_reduction(
1001 &execution,
1002 ¶meters,
1003 &prepared,
1004 &source,
1005 ReductionPlan::weighted_positive_real(),
1006 )
1007 .expect("prepared dataset should evaluate and reduce");
1008 assert_eq!(actual, expected);
1009 assert_eq!(sums, [7.5, 10.5]);
1010 let mut visited = vec![Vec::new(), Vec::new()];
1011 let parameter_sets = [
1012 (¶meters[0], "central"),
1013 (¶meters[1], "ensemble draw 0"),
1014 ];
1015 let sums = prepared_model
1016 .visit_prepared_many(
1017 &execution,
1018 ¶meter_sets,
1019 &prepared,
1020 &source,
1021 Some(ReductionPlan::weighted_positive_real()),
1022 |_, parameter_index, values| {
1023 visited[parameter_index].extend_from_slice(values);
1024 Ok(())
1025 },
1026 )
1027 .expect("prepared blocks should be visited and reduced");
1028 assert_eq!(visited, expected);
1029 assert_eq!(sums, [7.5, 10.5]);
1030 }
1031 }
1032
1033 #[test]
1034 fn prepared_block_visitors_run_inside_the_execution_owned_pool() {
1035 let model =
1036 CompiledModel::from_expr(&event_scalar("x")).expect("test model should compile");
1037 let execution = Execution::local(ExecutionOptions {
1038 device: Device::Cpu(CpuOptions {
1039 threads: ThreadPolicy::Fixed(2),
1040 jit: JitPolicy::Disabled,
1041 }),
1042 ..ExecutionOptions::default()
1043 })
1044 .expect("fixed-thread execution should build");
1045 let prepared_model =
1046 PreparedModel::prepare(&model, &execution).expect("model should prepare");
1047 let source = dataset(false);
1048 let prepared = prepared_model
1049 .prepare_dataset(&execution, &source)
1050 .expect("dataset should prepare");
1051 let parameters = model.params().values(&[]).expect("parameters should build");
1052 let mut visitor_pool_threads = Vec::new();
1053
1054 prepared_model
1055 .visit_prepared_many_parallel(
1056 &execution,
1057 &[(¶meters, "central")],
1058 &prepared,
1059 &source,
1060 None,
1061 |_, _, _| {
1062 visitor_pool_threads.push(rayon::current_num_threads());
1063 Ok(())
1064 },
1065 )
1066 .expect("prepared blocks should be visited");
1067
1068 assert!(!visitor_pool_threads.is_empty());
1069 assert!(visitor_pool_threads.iter().all(|&threads| threads == 2));
1070 }
1071
1072 #[test]
1073 fn existing_prepared_visitor_accepts_non_send_callbacks() {
1074 let model =
1075 CompiledModel::from_expr(&event_scalar("x")).expect("test model should compile");
1076 let execution = Execution::default();
1077 let prepared_model =
1078 PreparedModel::prepare(&model, &execution).expect("model should prepare");
1079 let source = dataset(false);
1080 let prepared = prepared_model
1081 .prepare_dataset(&execution, &source)
1082 .expect("dataset should prepare");
1083 let parameters = model.params().values(&[]).expect("parameters should build");
1084 let visited = Rc::new(RefCell::new(Vec::new()));
1085 let captured = Rc::clone(&visited);
1086
1087 prepared_model
1088 .visit_prepared_many(
1089 &execution,
1090 &[(¶meters, "central")],
1091 &prepared,
1092 &source,
1093 None,
1094 move |_, _, values| {
1095 captured.borrow_mut().extend_from_slice(values);
1096 Ok(())
1097 },
1098 )
1099 .expect("non-Send visitor should remain supported");
1100
1101 assert_eq!(visited.borrow().len(), 3);
1102 }
1103}
1104
1105#[cfg(all(test, feature = "wgpu"))]
1106mod tests {
1107 use std::sync::Arc;
1108
1109 use laddu_compile::{CompiledModel, ReductionPlan};
1110 use laddu_data::{
1111 data::{Dataset, EventBatch, OwnedEvent},
1112 schema::Schema,
1113 };
1114 use laddu_expr::{complex, event_scalar, parameter};
1115
1116 use super::*;
1117 use crate::{CpuOptions, Device, ExecutionOptions, GpuBackend, GpuOptions, Precision};
1118
1119 #[test]
1120 fn wgpu_dataset_plan_resolves_storage_and_chunk_limits_without_hardware() {
1121 let mut read_plan = laddu_data::io::ReadPlan::serial();
1122 read_plan.chunk_size = Some(3);
1123 let plan = WgpuDatasetPlan::resolve(
1124 read_plan,
1125 MemoryPolicy::Fastest,
1126 100,
1127 MemoryFootprint::new(100, 8),
1128 MemoryFootprint::new(200, 4),
1129 1_000,
1130 Some(10_000),
1131 )
1132 .unwrap();
1133
1134 assert_eq!(plan.preparation_plan.storage(), CacheStorage::Resident);
1135 assert_eq!(plan.read_plan.chunk_size, Some(3));
1136 assert_eq!(plan.device_decision.chunk_events, 100);
1137 assert_eq!(plan.device_decision.estimated_peak_bytes, 600);
1138 }
1139
1140 #[test]
1141 fn wgpu_dataset_plan_falls_back_to_streaming_when_resident_does_not_fit() {
1142 let plan = WgpuDatasetPlan::resolve(
1143 laddu_data::io::ReadPlan::serial(),
1144 MemoryPolicy::Fastest,
1145 100,
1146 MemoryFootprint::new(100, 8),
1147 MemoryFootprint::new(500, 10),
1148 1_000,
1149 Some(1_000),
1150 )
1151 .unwrap();
1152
1153 assert_eq!(plan.preparation_plan.storage(), CacheStorage::Streaming);
1154 assert_eq!(plan.read_plan.chunk_size, Some(50));
1155 assert_eq!(plan.device_decision.chunk_events, 50);
1156 assert_eq!(plan.device_decision.estimated_peak_bytes, 1_000);
1157 }
1158
1159 #[test]
1160 fn wgpu_dataset_plan_reports_host_failure_before_missing_device_pool() {
1161 let error = WgpuDatasetPlan::resolve(
1162 laddu_data::io::ReadPlan::serial(),
1163 MemoryPolicy::Fastest,
1164 1,
1165 MemoryFootprint::new(100, 8),
1166 MemoryFootprint::new(200, 4),
1167 100,
1168 None,
1169 )
1170 .unwrap_err();
1171
1172 assert!(matches!(
1173 error,
1174 RuntimeError::Memory(laddu_memory::MemoryError::BudgetExceeded { resource, .. })
1175 if resource == "WGPU host staging"
1176 ));
1177 }
1178
1179 #[test]
1180 #[ignore = "requires a WGPU-compatible hardware adapter"]
1181 fn wgpu_resident_and_streaming_reductions_match_f32_cpu() {
1182 let scale = laddu_expr::Expr::from(parameter!("scale", initial: 1.25));
1183 let offset = laddu_expr::Expr::from(parameter!("offset", initial: 0.5));
1184 let x = event_scalar("x");
1185 let expression = (x.clone() * scale.clone() + offset.clone()).sin()
1186 + complex(scale, offset).norm_sqr()
1187 + 2.0;
1188 let model = CompiledModel::from_expr(&expression).unwrap();
1189 let params = model.params().default_values();
1190 let schema = Arc::new(Schema::new(std::iter::empty::<&str>(), ["x"], true).unwrap());
1191 let dataset = Dataset::from_batches(vec![
1192 EventBatch::from_events(
1193 schema.clone(),
1194 [
1195 OwnedEvent::weighted(vec![], vec![0.25], 0.5),
1196 OwnedEvent::weighted(vec![], vec![0.75], 1.5),
1197 ],
1198 )
1199 .unwrap(),
1200 EventBatch::from_events(schema, [OwnedEvent::weighted(vec![], vec![1.25], 2.0)])
1201 .unwrap(),
1202 ])
1203 .unwrap();
1204 let wgpu_execution = Execution::local(ExecutionOptions {
1205 device: Device::Gpu(GpuOptions {
1206 backend: GpuBackend::Wgpu,
1207 ..GpuOptions::default()
1208 }),
1209 memory: crate::MemoryPlan::host_device(
1210 crate::MemoryBudget::Auto,
1211 crate::MemoryBudget::Bytes(256),
1212 ),
1213 precision: Precision::F32,
1214 ..ExecutionOptions::default()
1215 })
1216 .unwrap();
1217 let cpu_execution = Execution::local(ExecutionOptions {
1218 device: Device::Cpu(CpuOptions::default()),
1219 precision: Precision::F32,
1220 ..ExecutionOptions::default()
1221 })
1222 .unwrap();
1223 let wgpu = PreparedModel::prepare(&model, &wgpu_execution).unwrap();
1224 let cpu = PreparedModel::prepare(&model, &cpu_execution).unwrap();
1225 let resident = wgpu
1226 .prepare_dataset(&wgpu_execution, &dataset.clone().resident())
1227 .unwrap();
1228 let streaming = wgpu
1229 .prepare_dataset(&wgpu_execution, &dataset.clone().streaming())
1230 .unwrap();
1231 let cpu_data = cpu.prepare_dataset(&cpu_execution, &dataset).unwrap();
1232
1233 assert_eq!(resident.stats().storage(), CacheStorage::Resident);
1234 assert_eq!(streaming.stats().storage(), CacheStorage::Streaming);
1235 assert!(resident.stats().resident_bytes() > 0);
1236 assert_eq!(streaming.stats().resident_bytes(), 0);
1237
1238 let cpu_reduction = cpu
1239 .reduce_with_gradient(
1240 &cpu_execution,
1241 ¶ms,
1242 &cpu_data,
1243 ReductionPlan::weighted_real(),
1244 )
1245 .unwrap();
1246 let resident_reduction = wgpu
1247 .reduce_with_gradient(
1248 &wgpu_execution,
1249 ¶ms,
1250 &resident,
1251 ReductionPlan::weighted_real(),
1252 )
1253 .unwrap();
1254 let streaming_reduction = wgpu
1255 .reduce_with_gradient(
1256 &wgpu_execution,
1257 ¶ms,
1258 &streaming,
1259 ReductionPlan::weighted_real(),
1260 )
1261 .unwrap();
1262
1263 for actual in [&resident_reduction, &streaming_reduction] {
1264 assert!((actual.value() - cpu_reduction.value()).abs() <= 1.0e-4);
1265 assert_eq!(actual.gradient().len(), cpu_reduction.gradient().len());
1266 for (actual, expected) in actual.gradient().iter().zip(cpu_reduction.gradient()) {
1267 assert!((actual - expected).abs() <= 1.0e-4);
1268 }
1269 }
1270 assert!((resident_reduction.value() - streaming_reduction.value()).abs() <= 1.0e-6);
1271 for (resident, streaming) in resident_reduction
1272 .gradient()
1273 .iter()
1274 .zip(streaming_reduction.gradient())
1275 {
1276 assert!((resident - streaming).abs() <= 1.0e-6);
1277 }
1278 }
1279}