Skip to main content

laddu_runtime/
backend.rs

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/// A compiled model prepared for a concrete execution backend.
24#[derive(Clone, Debug)]
25pub enum PreparedModel {
26    /// A model prepared for CPU execution.
27    Cpu(Arc<CpuPlan>),
28    #[cfg(feature = "wgpu")]
29    /// A model prepared for WebGPU execution.
30    Wgpu(WgpuPlan),
31}
32
33/// A dataset prepared for a concrete execution backend.
34#[derive(Clone, Debug)]
35pub enum PreparedDataset {
36    /// A dataset prepared for CPU execution.
37    Cpu(CpuPreparedDataset),
38    #[cfg(feature = "wgpu")]
39    /// A dataset prepared for WebGPU execution.
40    Wgpu(WgpuPreparedDataset),
41}
42
43impl PreparedDataset {
44    /// Returns statistics collected while preparing the dataset.
45    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    /// Returns the event-scalar columns required by this prepared model.
56    ///
57    /// Requirements are deduplicated while retaining compiled graph order.
58    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    /// Estimate peak host workspace when visiting one batch across parameter sets.
67    #[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        // Evaluation can hold parallel block outputs and their flattened result.
75        cache.saturating_add(events.saturating_mul(64))
76    }
77
78    /// Visit parameter sets while a single batch's event cache is active.
79    /// The caller reserves batch workspace using `batch_memory_estimate`.
80    ///
81    /// # Errors
82    /// Returns an error for incompatible parameters/events or evaluation failure.
83    #[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    /// Evaluates the model for every event in a batch.
118    ///
119    /// # Errors
120    ///
121    /// Returns [`RuntimeError`] when parameters or event columns are
122    /// incompatible, evaluation fails, or a matrix solve is singular.
123    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    /// Evaluates a vector-root model and returns one output column per root
145    /// element. CPU supports this query-oriented ABI; scalar WGPU kernels
146    /// retain their existing single-output contract.
147    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    /// Evaluates the model and its free-parameter gradient for every event in a batch.
163    ///
164    /// # Errors
165    ///
166    /// Returns [`RuntimeError`] when inputs are incompatible, differentiation
167    /// or evaluation fails, or the selected backend lacks event-wise gradients.
168    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    /// Prepares a compiled model for the supplied execution context.
183    ///
184    /// # Errors
185    ///
186    /// Returns [`RuntimeError`] when model lowering, differentiation, backend
187    /// initialization, or precision selection fails.
188    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    /// Prepares a dataset for repeated evaluation with this model.
206    ///
207    /// # Errors
208    ///
209    /// Returns [`RuntimeError`] when the dataset cannot be read or cached, its
210    /// schema is incompatible, or backend preparation fails.
211    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    /// Evaluates every event in a prepared dataset while preserving source order.
228    ///
229    /// Backends with a prepared event adapter reuse their retained or streaming
230    /// prepared blocks. Other backends evaluate the supplied source through the
231    /// already-selected backend; this operation never substitutes a backend.
232    ///
233    /// # Errors
234    ///
235    /// Returns [`RuntimeError`] when model and dataset backends differ, source
236    /// streaming fails, or event evaluation fails.
237    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    /// Evaluates multiple parameter sets while each prepared event block is active.
250    ///
251    /// Output rows retain parameter-set order and each row retains source-event order.
252    ///
253    /// # Errors
254    ///
255    /// Returns [`RuntimeError`] when model and dataset backends differ, source
256    /// streaming fails, or event evaluation fails.
257    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    /// Visits one bounded block of prepared values at a time for several
284    /// parameter sets without retaining full-dataset value rows.
285    #[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    /// Visits prepared CPU values on the execution-owned thread pool.
322    #[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    /// Evaluates and reduces multiple parameter sets while each prepared block is active.
361    ///
362    /// # Errors
363    ///
364    /// Returns [`RuntimeError`] when model and dataset backends differ, source
365    /// streaming fails, event evaluation fails, or the reduction rejects a value.
366    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    /// Executes a weighted scalar reduction over a prepared dataset.
396    ///
397    /// # Errors
398    ///
399    /// Returns [`RuntimeError`] when model and dataset backends differ, inputs
400    /// are incompatible, evaluation fails, or the reduction domain is invalid.
401    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    /// Executes a weighted reduction and computes its free-parameter gradient.
425    ///
426    /// # Errors
427    ///
428    /// Returns [`RuntimeError`] when model and dataset backends differ,
429    /// differentiation or evaluation fails, or the reduction domain is invalid.
430    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/// A compiled model prepared for WebGPU execution.
561#[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/// Dataset storage prepared for WebGPU evaluation.
667#[derive(Clone)]
668pub enum WgpuPreparedDataset {
669    /// GPU-resident prepared batches.
670    Resident {
671        /// Prepared GPU batches.
672        batches: std::sync::Arc<[laddu_wgpu::WgpuPreparedBatch]>,
673        /// Preparation statistics.
674        stats: PreparedDatasetStats,
675        /// Persistent device-memory reservation.
676        memory_lease: crate::MemoryLease,
677    },
678    /// Source data streamed and prepared one batch at a time.
679    Streaming {
680        /// Source dataset.
681        dataset: Dataset,
682        /// Read plan used on each pass.
683        read_plan: laddu_data::io::ReadPlan,
684        /// Reusable prepared-batch workspace.
685        workspace: std::sync::Arc<std::sync::Mutex<Option<laddu_wgpu::WgpuPreparedBatch>>>,
686        /// Preparation statistics.
687        stats: PreparedDatasetStats,
688        /// Peak transient device bytes reserved during reductions.
689        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    /// Returns statistics collected while preparing the dataset.
706    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, &parameters, &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                    &parameters,
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                (&parameters[0], "central"),
1013                (&parameters[1], "ensemble draw 0"),
1014            ];
1015            let sums = prepared_model
1016                .visit_prepared_many(
1017                    &execution,
1018                    &parameter_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                &[(&parameters, "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                &[(&parameters, "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                &params,
1242                &cpu_data,
1243                ReductionPlan::weighted_real(),
1244            )
1245            .unwrap();
1246        let resident_reduction = wgpu
1247            .reduce_with_gradient(
1248                &wgpu_execution,
1249                &params,
1250                &resident,
1251                ReductionPlan::weighted_real(),
1252            )
1253            .unwrap();
1254        let streaming_reduction = wgpu
1255            .reduce_with_gradient(
1256                &wgpu_execution,
1257                &params,
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}