Skip to main content

laddu_runtime/cpu/
reduction.rs

1use laddu_compile::{ReductionPlan, ReductionTransform};
2#[cfg(test)]
3use laddu_data::data::accurate::AccurateComplex64;
4use laddu_data::data::accurate::AccurateF64;
5use laddu_data::{LadduDataResult, data::EventBatch};
6use laddu_expr::parameters::ParamValues;
7use laddu_memory::MemoryLease;
8use num::complex::Complex64;
9use rayon::prelude::*;
10
11use crate::execution::Execution;
12#[cfg(feature = "jit")]
13use crate::jit::{JitGradientKernel, JitScalarKernel};
14
15use super::{
16    CpuCachedBatch, CpuCachedDataset, CpuPlan, CpuPreparedDataset, F32KernelInput, Precision,
17    ReductionEvaluation, RuntimeError, RuntimeResult, SCALAR_BLOCK_SIZE, ScalarEventWorkspace,
18    ValueGradient,
19};
20
21struct RealGradientAccumulator {
22    value: AccurateF64,
23    gradient: Vec<AccurateF64>,
24}
25
26/// The common input boundary for all CPU reductions.
27///
28/// Resident batches are borrowed, while streaming batches are cached as they
29/// are pulled. Keeping that distinction private avoids rebuilding a temporary
30/// one-batch dataset merely to select a reduction implementation.
31struct PreparedBatchStream<'a> {
32    source: PreparedBatchSource<'a>,
33    plan: &'a CpuPlan,
34}
35
36enum PreparedBatchSource<'a> {
37    Resident(std::slice::Iter<'a, CpuCachedBatch>),
38    Streaming {
39        batches: Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send + 'a>,
40        _memory: MemoryLease,
41    },
42}
43
44enum PreparedBatch<'a> {
45    Borrowed(&'a CpuCachedBatch),
46    Owned(CpuCachedBatch),
47}
48
49type PreparedManyEvaluation = (Vec<Vec<Complex64>>, Option<Vec<f64>>);
50
51impl<'a> PreparedBatch<'a> {
52    fn cached(&self) -> &CpuCachedBatch {
53        match self {
54            Self::Borrowed(batch) => batch,
55            Self::Owned(batch) => batch,
56        }
57    }
58}
59
60impl<'a> PreparedBatchStream<'a> {
61    fn prepare(
62        plan: &'a CpuPlan,
63        execution: &Execution,
64        dataset: &'a CpuPreparedDataset,
65    ) -> RuntimeResult<Self> {
66        let source = match dataset {
67            CpuPreparedDataset::Resident { dataset, .. } => {
68                PreparedBatchSource::Resident(dataset.batches().iter())
69            }
70            CpuPreparedDataset::Streaming {
71                dataset,
72                read_plan,
73                transient_bytes,
74                ..
75            } => {
76                let memory = execution
77                    .host_memory()
78                    .reserve(*transient_bytes)
79                    .map_err(RuntimeError::from)?;
80                let batches = dataset
81                    .stream_with_plan(*read_plan)
82                    .map_err(|error| RuntimeError::Data(error.to_string()))?;
83                PreparedBatchSource::Streaming {
84                    batches,
85                    _memory: memory,
86                }
87            }
88        };
89        Ok(Self { source, plan })
90    }
91
92    fn next(&mut self) -> RuntimeResult<Option<PreparedBatch<'a>>> {
93        match &mut self.source {
94            PreparedBatchSource::Resident(batches) => {
95                Ok(batches.next().map(PreparedBatch::Borrowed))
96            }
97            PreparedBatchSource::Streaming { batches, .. } => batches
98                .next()
99                .transpose()
100                .map_err(|error| RuntimeError::Data(error.to_string()))?
101                .map(|batch| {
102                    self.plan
103                        .cache_event_batch(&batch)
104                        .map(|cache| PreparedBatch::Owned(CpuCachedBatch { cache }))
105                })
106                .transpose(),
107        }
108    }
109}
110
111struct ValueReducer {
112    total: AccurateF64,
113}
114
115impl ValueReducer {
116    fn new() -> Self {
117        Self {
118            total: AccurateF64::zero(),
119        }
120    }
121
122    fn consume(
123        &mut self,
124        plan: &CpuPlan,
125        execution: &Execution,
126        params: &ParamValues,
127        batch: &CpuCachedBatch,
128        reduction: ReductionPlan,
129    ) -> RuntimeResult<()> {
130        let value = if execution.is_parallel() && batch.len().div_ceil(SCALAR_BLOCK_SIZE) >= 2 {
131            execution.install(|| {
132                plan.par_try_weighted_sum_batch(params, batch, |value| {
133                    plan.apply_reduction(reduction, value)
134                })
135            })?
136        } else {
137            plan.try_weighted_sum_batch(params, batch, |value| {
138                plan.apply_reduction(reduction, value)
139            })?
140        };
141        self.total.push(value);
142        Ok(())
143    }
144
145    fn finish(self) -> f64 {
146        self.total.finish()
147    }
148}
149
150struct GradientReducer {
151    total: RealGradientAccumulator,
152}
153
154impl GradientReducer {
155    fn new(parameter_count: usize) -> Self {
156        Self {
157            total: RealGradientAccumulator::zero(parameter_count),
158        }
159    }
160
161    fn consume<F>(
162        &mut self,
163        plan: &CpuPlan,
164        execution: &Execution,
165        params: &ParamValues,
166        batch: &CpuCachedBatch,
167        transform: &F,
168    ) -> RuntimeResult<()>
169    where
170        F: Fn(Complex64) -> RuntimeResult<(f64, f64)> + Send + Sync,
171    {
172        let mut transform = transform;
173        let (value, gradient) =
174            if execution.is_parallel() && batch.len().div_ceil(SCALAR_BLOCK_SIZE) >= 2 {
175                execution.install(|| {
176                    plan.par_try_weighted_real_sum_with_gradient_batch(params, batch, &transform)
177                })?
178            } else {
179                plan.try_weighted_real_sum_with_gradient_batch(params, batch, &mut transform)?
180            };
181        self.total.value.push(value);
182        for (sum, partial) in self.total.gradient.iter_mut().zip(gradient) {
183            sum.push(partial);
184        }
185        Ok(())
186    }
187
188    fn finish(self) -> (f64, Vec<f64>) {
189        self.total.finish()
190    }
191}
192
193impl RealGradientAccumulator {
194    fn zero(parameter_count: usize) -> Self {
195        Self {
196            value: AccurateF64::zero(),
197            gradient: (0..parameter_count).map(|_| AccurateF64::zero()).collect(),
198        }
199    }
200
201    fn push(&mut self, weight: f64, value: f64, derivative: f64, model_gradient: &[Complex64]) {
202        self.value.push(weight * value);
203        for (sum, model_derivative) in self.gradient.iter_mut().zip(model_gradient) {
204            sum.push(weight * derivative * model_derivative.re);
205        }
206    }
207
208    fn push_f32(&mut self, weight: f64, value: f64, derivative: f64, model_gradient: &[f32]) {
209        self.value.push(weight * value);
210        for (sum, model_derivative) in self.gradient.iter_mut().zip(model_gradient) {
211            sum.push(weight * derivative * f64::from(*model_derivative));
212        }
213    }
214
215    fn merge(&mut self, other: Self) {
216        self.value.merge(other.value);
217        for (target, source) in self.gradient.iter_mut().zip(other.gradient) {
218            target.merge(source);
219        }
220    }
221
222    fn finish(self) -> (f64, Vec<f64>) {
223        (
224            self.value.finish(),
225            self.gradient.into_iter().map(AccurateF64::finish).collect(),
226        )
227    }
228}
229
230impl CpuPlan {
231    pub(crate) fn visit_prepared_dataset_many<F>(
232        &self,
233        execution: &Execution,
234        parameter_sets: &[(&ParamValues, &str)],
235        dataset: &CpuPreparedDataset,
236        reduction: Option<ReductionPlan>,
237        pool_installed: bool,
238        mut consume: F,
239    ) -> RuntimeResult<Vec<f64>>
240    where
241        F: FnMut(usize, usize, &[Complex64]) -> RuntimeResult<()>,
242    {
243        let local = (|| {
244            let mut stream = PreparedBatchStream::prepare(self, execution, dataset)?;
245            let mut offset = 0;
246            let mut sums = reduction.map(|_| {
247                parameter_sets
248                    .iter()
249                    .map(|_| AccurateF64::zero())
250                    .collect::<Vec<_>>()
251            });
252            while let Some(batch) = stream.next()? {
253                let batch = batch.cached();
254                for (index, &(parameters, context)) in parameter_sets.iter().enumerate() {
255                    let values = self
256                        .evaluate_prepared_batch(execution, parameters, batch, pool_installed)
257                        .map_err(|error| {
258                            RuntimeError::Parameter(format!("{context} evaluation failed: {error}"))
259                        })?;
260                    if let (Some(reduction), Some(sums)) = (reduction, sums.as_mut()) {
261                        for (weight, value) in batch.weights().iter().zip(&values) {
262                            let reduced = reduction.apply(*value).map_err(|error| {
263                                RuntimeError::Parameter(format!(
264                                    "{context} reduction failed: {error}"
265                                ))
266                            })?;
267                            sums[index].push(*weight * reduced.value());
268                        }
269                    }
270                    consume(offset, index, &values)?;
271                }
272                offset += batch.weights().len();
273            }
274            Ok(sums
275                .unwrap_or_default()
276                .into_iter()
277                .map(AccurateF64::finish)
278                .collect::<Vec<_>>())
279        })();
280        if !execution.all_succeeded(local.is_ok()) {
281            return local.and(Err(RuntimeError::DistributedPeerFailure));
282        }
283        Ok(local?
284            .into_iter()
285            .map(|sum| execution.sum_f64(sum))
286            .collect())
287    }
288
289    pub(crate) fn evaluate_prepared_batch(
290        &self,
291        execution: &Execution,
292        params: &ParamValues,
293        batch: &CpuCachedBatch,
294        pool_installed: bool,
295    ) -> RuntimeResult<Vec<Complex64>> {
296        let block_count = batch.len().div_ceil(SCALAR_BLOCK_SIZE);
297        if !execution.is_parallel() || block_count < 2 {
298            return self.evaluate_cache(params, batch.cache());
299        }
300        self.check_batch_cache(batch.cache())?;
301        let invariant = self.scalar_invariant_values(params)?;
302        #[cfg(feature = "jit")]
303        let jit_cache = self
304            .scalar_jit_kernel()
305            .map(|_| JitScalarKernel::prepare_cache(batch.cache()));
306        let evaluate = || {
307            (0..block_count)
308                .into_par_iter()
309                .map(|block| {
310                    let start = block * SCALAR_BLOCK_SIZE;
311                    let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
312                    let mut workspace = ScalarEventWorkspace::default();
313                    let mut output = Vec::with_capacity(end - start);
314                    self.evaluate_cache_block_prepared(
315                        params,
316                        batch.cache(),
317                        start,
318                        end,
319                        invariant.as_ref(),
320                        &mut workspace,
321                        &mut output,
322                        #[cfg(feature = "jit")]
323                        jit_cache.as_ref(),
324                    )?;
325                    Ok(output)
326                })
327                .collect::<RuntimeResult<Vec<_>>>()
328        };
329        let blocks = if pool_installed {
330            evaluate()?
331        } else {
332            execution.install(evaluate)?
333        };
334        Ok(blocks.into_iter().flatten().collect())
335    }
336
337    pub(crate) fn evaluate_prepared_dataset_many(
338        &self,
339        execution: &Execution,
340        params: &[ParamValues],
341        dataset: &CpuPreparedDataset,
342    ) -> RuntimeResult<Vec<Vec<Complex64>>> {
343        self.evaluate_prepared_dataset_many_local(execution, params, dataset, None)
344            .map(|(values, _)| values)
345    }
346
347    pub(crate) fn evaluate_prepared_dataset_many_with_reduction(
348        &self,
349        execution: &Execution,
350        params: &[ParamValues],
351        dataset: &CpuPreparedDataset,
352        reduction: ReductionPlan,
353    ) -> RuntimeResult<(Vec<Vec<Complex64>>, Vec<f64>)> {
354        let local =
355            self.evaluate_prepared_dataset_many_local(execution, params, dataset, Some(reduction));
356        if !execution.all_succeeded(local.is_ok()) {
357            return local.and(Err(RuntimeError::DistributedPeerFailure));
358        }
359        let (output, sums) = local?;
360        Ok((
361            output,
362            sums.unwrap_or_default()
363                .into_iter()
364                .map(|sum| execution.sum_f64(sum))
365                .collect(),
366        ))
367    }
368
369    fn evaluate_prepared_dataset_many_local(
370        &self,
371        execution: &Execution,
372        params: &[ParamValues],
373        dataset: &CpuPreparedDataset,
374        reduction: Option<ReductionPlan>,
375    ) -> RuntimeResult<PreparedManyEvaluation> {
376        let mut stream = PreparedBatchStream::prepare(self, execution, dataset)?;
377        let mut output = params
378            .iter()
379            .map(|_| Vec::with_capacity(dataset.stats().local_events()))
380            .collect::<Vec<_>>();
381        let mut sums = reduction.map(|_| {
382            params
383                .iter()
384                .map(|_| AccurateF64::zero())
385                .collect::<Vec<_>>()
386        });
387        while let Some(batch) = stream.next()? {
388            let batch = batch.cached();
389            for (index, (parameters, values)) in params.iter().zip(&mut output).enumerate() {
390                let batch_values = self.evaluate_cache(parameters, batch.cache())?;
391                if let (Some(reduction), Some(sums)) = (reduction, sums.as_mut()) {
392                    for (weight, value) in batch.weights().iter().zip(&batch_values) {
393                        sums[index].push(*weight * reduction.apply(*value)?.value());
394                    }
395                }
396                values.extend(batch_values);
397            }
398        }
399        Ok((
400            output,
401            sums.map(|sums| sums.into_iter().map(AccurateF64::finish).collect()),
402        ))
403    }
404
405    /// Execute a weighted reduction over a prepared dataset.
406    ///
407    /// # Errors
408    ///
409    /// Returns [`RuntimeError`] when streaming, cache validation, evaluation,
410    /// or reduction fails, or another distributed worker reports failure.
411    pub fn reduce(
412        &self,
413        execution: &Execution,
414        params: &ParamValues,
415        dataset: &CpuPreparedDataset,
416        reduction: ReductionPlan,
417    ) -> RuntimeResult<f64> {
418        let local = (|| {
419            let mut stream = PreparedBatchStream::prepare(self, execution, dataset)?;
420            let mut reducer = ValueReducer::new();
421            while let Some(batch) = stream.next()? {
422                reducer.consume(self, execution, params, batch.cached(), reduction)?;
423            }
424            Ok(reducer.finish())
425        })();
426        if !execution.all_succeeded(local.is_ok()) {
427            return local.and(Err(RuntimeError::DistributedPeerFailure));
428        }
429        Ok(execution.sum_f64(local?))
430    }
431
432    /// Execute a weighted reduction and its free-parameter gradient.
433    ///
434    /// # Errors
435    ///
436    /// Returns [`RuntimeError`] when streaming, cache validation,
437    /// differentiation, evaluation, or reduction fails, or another distributed
438    /// worker reports failure.
439    pub fn reduce_with_gradient(
440        &self,
441        execution: &Execution,
442        params: &ParamValues,
443        dataset: &CpuPreparedDataset,
444        reduction: ReductionPlan,
445    ) -> RuntimeResult<ReductionEvaluation> {
446        let local = (|| {
447            let mut stream = PreparedBatchStream::prepare(self, execution, dataset)?;
448            let mut reducer = GradientReducer::new(self.free_parameter_count());
449            let transform = |value| {
450                reduction
451                    .apply(value)
452                    .map(|output| output.into_parts())
453                    .map_err(RuntimeError::from)
454            };
455            while let Some(batch) = stream.next()? {
456                reducer.consume(self, execution, params, batch.cached(), &transform)?;
457            }
458            Ok(reducer.finish())
459        })();
460        if !execution.all_succeeded(local.is_ok()) {
461            return local.and(Err(RuntimeError::DistributedPeerFailure));
462        }
463        let (value, gradient) = local?;
464        let value = execution.sum_f64(value);
465        let gradient = execution.sum_slice(&gradient);
466        Ok(ReductionEvaluation { value, gradient })
467    }
468
469    /// Evaluates every event in a fully cached dataset.
470    ///
471    /// # Errors
472    ///
473    /// Returns [`RuntimeError`] when parameters or a cache layout are
474    /// incompatible, evaluation fails, or a matrix is singular.
475    pub fn evaluate_cached_dataset(
476        &self,
477        params: &ParamValues,
478        dataset: &CpuCachedDataset,
479    ) -> RuntimeResult<Vec<Complex64>> {
480        let total_len = dataset.batches.iter().map(CpuCachedBatch::len).sum();
481        let mut out = Vec::with_capacity(total_len);
482        let invariant = self.scalar_invariant_values(params)?;
483        let mut workspace = ScalarEventWorkspace::default();
484        for batch in &dataset.batches {
485            self.check_batch_cache(batch.cache())?;
486            for row in 0..batch.len() {
487                out.push(self.evaluate_cache_row_prepared(
488                    params,
489                    batch.cache(),
490                    row,
491                    invariant.as_ref(),
492                    &mut workspace,
493                )?);
494            }
495        }
496        Ok(out)
497    }
498
499    /// Evaluates every event and gradient in a fully cached dataset.
500    ///
501    /// # Errors
502    ///
503    /// Returns [`RuntimeError`] when parameters or a cache layout are
504    /// incompatible, or differentiation or evaluation fails.
505    pub fn evaluate_cached_dataset_with_gradient(
506        &self,
507        params: &ParamValues,
508        dataset: &CpuCachedDataset,
509    ) -> RuntimeResult<Vec<ValueGradient>> {
510        let total_len = dataset.batches.iter().map(CpuCachedBatch::len).sum();
511        let mut out = Vec::with_capacity(total_len);
512        for batch in &dataset.batches {
513            out.extend(self.evaluate_cache_with_gradient(params, batch.cache())?);
514        }
515        Ok(out)
516    }
517
518    fn try_weighted_sum_batch<E, F>(
519        &self,
520        params: &ParamValues,
521        batch: &CpuCachedBatch,
522        mut f: F,
523    ) -> Result<f64, E>
524    where
525        E: From<RuntimeError>,
526        F: FnMut(Complex64) -> Result<f64, E>,
527    {
528        self.check_batch_cache(batch.cache())?;
529        let invariant = self.scalar_invariant_values(params)?;
530        let mut workspace = ScalarEventWorkspace::default();
531        let mut output = Vec::with_capacity(SCALAR_BLOCK_SIZE);
532        #[cfg(feature = "jit")]
533        let jit_cache = self
534            .scalar_jit_kernel()
535            .map(|_| JitScalarKernel::prepare_cache(batch.cache()));
536        let mut sum = AccurateF64::zero();
537        for start in (0..batch.len()).step_by(SCALAR_BLOCK_SIZE) {
538            let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
539            self.evaluate_cache_block_prepared(
540                params,
541                batch.cache(),
542                start,
543                end,
544                invariant.as_ref(),
545                &mut workspace,
546                &mut output,
547                #[cfg(feature = "jit")]
548                jit_cache.as_ref(),
549            )?;
550            for (lane, value) in output.iter().copied().enumerate() {
551                sum.push(batch.weights()[start + lane] * f(value)?);
552            }
553        }
554        Ok(sum.finish())
555    }
556
557    fn par_try_weighted_sum_batch<E, F>(
558        &self,
559        params: &ParamValues,
560        batch: &CpuCachedBatch,
561        f: F,
562    ) -> Result<f64, E>
563    where
564        E: From<RuntimeError> + Send,
565        F: Fn(Complex64) -> Result<f64, E> + Send + Sync,
566    {
567        self.check_batch_cache(batch.cache())?;
568        let invariant = self.scalar_invariant_values(params)?;
569        #[cfg(feature = "jit")]
570        let jit_cache = self
571            .scalar_jit_kernel()
572            .map(|_| JitScalarKernel::prepare_cache(batch.cache()));
573        let n_blocks = batch.len().div_ceil(SCALAR_BLOCK_SIZE);
574        let total = (0..n_blocks)
575            .into_par_iter()
576            .try_fold(
577                || {
578                    (
579                        AccurateF64::zero(),
580                        ScalarEventWorkspace::default(),
581                        Vec::new(),
582                    )
583                },
584                |(mut acc, mut workspace, mut output), block| {
585                    let start = block * SCALAR_BLOCK_SIZE;
586                    let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
587                    self.evaluate_cache_block_prepared(
588                        params,
589                        batch.cache(),
590                        start,
591                        end,
592                        invariant.as_ref(),
593                        &mut workspace,
594                        &mut output,
595                        #[cfg(feature = "jit")]
596                        jit_cache.as_ref(),
597                    )?;
598                    for (lane, value) in output.iter().copied().enumerate() {
599                        acc.push(batch.weights()[start + lane] * f(value)?);
600                    }
601                    Ok::<_, E>((acc, workspace, output))
602                },
603            )
604            .try_reduce(
605                || {
606                    (
607                        AccurateF64::zero(),
608                        ScalarEventWorkspace::default(),
609                        Vec::new(),
610                    )
611                },
612                |(mut lhs, workspace, output), (rhs, _, _)| {
613                    lhs.merge(rhs);
614                    Ok::<_, E>((lhs, workspace, output))
615                },
616            )?;
617        Ok(total.0.finish())
618    }
619
620    fn try_weighted_real_sum_with_gradient_batch<E, F>(
621        &self,
622        params: &ParamValues,
623        batch: &CpuCachedBatch,
624        transform: &mut F,
625    ) -> Result<(f64, Vec<f64>), E>
626    where
627        E: From<RuntimeError>,
628        F: FnMut(Complex64) -> Result<(f64, f64), E>,
629    {
630        self.check_batch_cache(batch.cache())?;
631        #[cfg(feature = "jit")]
632        if let (Some(value_kernel), Some(gradient_kernel)) =
633            (self.scalar_jit_kernel(), self.gradient_jit_kernel())
634        {
635            return self.try_weighted_real_sum_with_jit_gradient_batch(
636                params,
637                batch,
638                transform,
639                value_kernel,
640                gradient_kernel,
641            );
642        }
643        if self.precision != Precision::F32
644            && let Some(interpreter) = self.gradient_interpreter()
645            && let Some(mut state) = interpreter.prepare_real_blocks(params)?
646        {
647            let output_count = state.output_count();
648            let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
649            for block in 0..batch.len().div_ceil(SCALAR_BLOCK_SIZE) {
650                let start = block * SCALAR_BLOCK_SIZE;
651                let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
652                let outputs = state.evaluate(batch.cache(), start, end)?;
653                for (lane, row) in outputs.chunks_exact(output_count).enumerate() {
654                    let (value, derivative) = transform(row[0])?;
655                    total.push(batch.weights()[start + lane], value, derivative, &row[1..]);
656                }
657            }
658            return Ok(total.finish());
659        }
660        if self.precision == Precision::F32
661            && let Some(ir) = self.f32_gradient_fallback_real.as_ref()
662        {
663            let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
664            let mut gradient = Vec::new();
665            for row in 0..batch.len() {
666                let (value, model_gradient) = self.evaluate_f32_gradient_component_prepared(
667                    ir,
668                    params,
669                    F32KernelInput::Cache(Some((batch.cache(), row))),
670                    &mut gradient,
671                )?;
672                let (value, derivative) = transform(value)?;
673                total.push_f32(batch.weights()[row], value, derivative, model_gradient);
674            }
675            return Ok(total.finish());
676        }
677        let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
678        for row in 0..batch.len() {
679            let evaluation =
680                self.evaluate_cache_row_with_gradient_unchecked(params, batch.cache(), row)?;
681            let (value, derivative) = transform(evaluation.value())?;
682            total.push(
683                batch.weights()[row],
684                value,
685                derivative,
686                evaluation.gradient(),
687            );
688        }
689        Ok(total.finish())
690    }
691
692    fn par_try_weighted_real_sum_with_gradient_batch<E, F>(
693        &self,
694        params: &ParamValues,
695        batch: &CpuCachedBatch,
696        transform: &F,
697    ) -> Result<(f64, Vec<f64>), E>
698    where
699        E: From<RuntimeError> + Send,
700        F: Fn(Complex64) -> Result<(f64, f64), E> + Send + Sync,
701    {
702        #[cfg(feature = "jit")]
703        if let (Some(value_kernel), Some(gradient_kernel)) =
704            (self.scalar_jit_kernel(), self.gradient_jit_kernel())
705        {
706            return self.par_try_weighted_real_sum_with_jit_gradient_batch(
707                params,
708                batch,
709                transform,
710                value_kernel,
711                gradient_kernel,
712            );
713        }
714        if self.precision != Precision::F32
715            && let Some(interpreter) = self.gradient_interpreter()
716            && let Some(state) = interpreter.prepare_real_blocks(params)?
717        {
718            let output_count = state.output_count();
719            let partial = (0..batch.len().div_ceil(SCALAR_BLOCK_SIZE))
720                .into_par_iter()
721                .try_fold(
722                    || {
723                        (
724                            RealGradientAccumulator::zero(self.free_parameter_count()),
725                            state.clone(),
726                        )
727                    },
728                    |(mut accumulator, mut state), block| {
729                        let start = block * SCALAR_BLOCK_SIZE;
730                        let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
731                        let outputs = state.evaluate(batch.cache(), start, end)?;
732                        for (lane, row) in outputs.chunks_exact(output_count).enumerate() {
733                            let (value, derivative) = transform(row[0])?;
734                            accumulator.push(
735                                batch.weights()[start + lane],
736                                value,
737                                derivative,
738                                &row[1..],
739                            );
740                        }
741                        Ok::<_, E>((accumulator, state))
742                    },
743                )
744                .try_reduce(
745                    || {
746                        (
747                            RealGradientAccumulator::zero(self.free_parameter_count()),
748                            state.clone(),
749                        )
750                    },
751                    |(mut lhs, state), (rhs, _)| {
752                        lhs.merge(rhs);
753                        Ok::<_, E>((lhs, state))
754                    },
755                )?;
756            return Ok(partial.0.finish());
757        }
758        if self.precision == Precision::F32
759            && let Some(ir) = self.f32_gradient_fallback_real.as_ref()
760        {
761            let partial = (0..batch.len())
762                .into_par_iter()
763                .try_fold(
764                    || {
765                        (
766                            RealGradientAccumulator::zero(self.free_parameter_count()),
767                            Vec::new(),
768                        )
769                    },
770                    |(mut accumulator, mut gradient), row| {
771                        let (value, model_gradient) = self
772                            .evaluate_f32_gradient_component_prepared(
773                                ir,
774                                params,
775                                F32KernelInput::Cache(Some((batch.cache(), row))),
776                                &mut gradient,
777                            )?;
778                        let (value, derivative) = transform(value)?;
779                        accumulator.push_f32(
780                            batch.weights()[row],
781                            value,
782                            derivative,
783                            model_gradient,
784                        );
785                        Ok::<_, E>((accumulator, gradient))
786                    },
787                )
788                .try_reduce(
789                    || {
790                        (
791                            RealGradientAccumulator::zero(self.free_parameter_count()),
792                            Vec::new(),
793                        )
794                    },
795                    |(mut lhs, gradient), (rhs, _)| {
796                        lhs.merge(rhs);
797                        Ok::<_, E>((lhs, gradient))
798                    },
799                )?;
800            return Ok(partial.0.finish());
801        }
802        let partial = (0..batch.len())
803            .into_par_iter()
804            .try_fold(
805                || RealGradientAccumulator::zero(self.free_parameter_count()),
806                |mut accumulator, row| {
807                    let evaluation = self.evaluate_cache_row_with_gradient_unchecked(
808                        params,
809                        batch.cache(),
810                        row,
811                    )?;
812                    let (value, derivative) = transform(evaluation.value())?;
813                    accumulator.push(
814                        batch.weights()[row],
815                        value,
816                        derivative,
817                        evaluation.gradient(),
818                    );
819                    Ok::<_, E>(accumulator)
820                },
821            )
822            .try_reduce(
823                || RealGradientAccumulator::zero(self.free_parameter_count()),
824                |mut lhs, rhs| {
825                    lhs.merge(rhs);
826                    Ok::<_, E>(lhs)
827                },
828            )?;
829        Ok(partial.finish())
830    }
831
832    #[cfg(feature = "jit")]
833    fn try_weighted_real_sum_with_jit_gradient_batch<E, F>(
834        &self,
835        params: &ParamValues,
836        batch: &CpuCachedBatch,
837        transform: &mut F,
838        value_kernel: &JitScalarKernel,
839        gradient_kernel: &JitGradientKernel,
840    ) -> Result<(f64, Vec<f64>), E>
841    where
842        E: From<RuntimeError>,
843        F: FnMut(Complex64) -> Result<(f64, f64), E>,
844    {
845        let cache = JitScalarKernel::prepare_cache(batch.cache());
846        let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
847        let mut values = Vec::new();
848        let mut tangents = Vec::new();
849        let mut derivatives = Vec::new();
850        for block in 0..batch.len().div_ceil(SCALAR_BLOCK_SIZE) {
851            let start = block * SCALAR_BLOCK_SIZE;
852            let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
853            value_kernel.evaluate_prepared(params, &cache, start, end, &mut values)?;
854            derivatives.clear();
855            derivatives.reserve(values.len());
856            for (lane, value) in values.iter().copied().enumerate() {
857                let (value, derivative) = transform(value)?;
858                total.value.push(batch.weights()[start + lane] * value);
859                derivatives.push(batch.weights()[start + lane] * derivative);
860            }
861            gradient_kernel.evaluate_prepared(params, &cache, start, end, 0, &mut tangents)?;
862            for (lane, factor) in derivatives.iter().enumerate() {
863                for free_index in 0..self.free_parameter_count() {
864                    total.gradient[free_index]
865                        .push(factor * tangents[lane * self.free_parameter_count() + free_index]);
866                }
867            }
868        }
869        Ok(total.finish())
870    }
871
872    #[cfg(feature = "jit")]
873    fn par_try_weighted_real_sum_with_jit_gradient_batch<E, F>(
874        &self,
875        params: &ParamValues,
876        batch: &CpuCachedBatch,
877        transform: &F,
878        value_kernel: &JitScalarKernel,
879        gradient_kernel: &JitGradientKernel,
880    ) -> Result<(f64, Vec<f64>), E>
881    where
882        E: From<RuntimeError> + Send,
883        F: Fn(Complex64) -> Result<(f64, f64), E> + Send + Sync,
884    {
885        let cache = JitScalarKernel::prepare_cache(batch.cache());
886        let partial = (0..batch.len().div_ceil(SCALAR_BLOCK_SIZE))
887            .into_par_iter()
888            .try_fold(
889                || {
890                    (
891                        RealGradientAccumulator::zero(self.free_parameter_count()),
892                        Vec::new(),
893                        Vec::new(),
894                        Vec::new(),
895                    )
896                },
897                |(mut accumulator, mut values, mut tangents, mut derivatives), block| {
898                    let start = block * SCALAR_BLOCK_SIZE;
899                    let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
900                    value_kernel.evaluate_prepared(params, &cache, start, end, &mut values)?;
901                    derivatives.clear();
902                    for (lane, value) in values.iter().copied().enumerate() {
903                        let (value, derivative) = transform(value)?;
904                        let weight = batch.weights()[start + lane];
905                        accumulator.value.push(weight * value);
906                        derivatives.push(weight * derivative);
907                    }
908                    gradient_kernel.evaluate_prepared(
909                        params,
910                        &cache,
911                        start,
912                        end,
913                        0,
914                        &mut tangents,
915                    )?;
916                    for (lane, factor) in derivatives.iter().enumerate() {
917                        for free_index in 0..self.free_parameter_count() {
918                            accumulator.gradient[free_index].push(
919                                factor * tangents[lane * self.free_parameter_count() + free_index],
920                            );
921                        }
922                    }
923                    Ok::<_, E>((accumulator, values, tangents, derivatives))
924                },
925            )
926            .try_reduce(
927                || {
928                    (
929                        RealGradientAccumulator::zero(self.free_parameter_count()),
930                        Vec::new(),
931                        Vec::new(),
932                        Vec::new(),
933                    )
934                },
935                |(mut lhs, values, tangents, derivatives), (rhs, _, _, _)| {
936                    lhs.merge(rhs);
937                    Ok::<_, E>((lhs, values, tangents, derivatives))
938                },
939            )?;
940        Ok(partial.0.finish())
941    }
942
943    #[cfg(test)]
944    fn try_weighted_sum_cached<E, F>(
945        &self,
946        params: &ParamValues,
947        dataset: &CpuCachedDataset,
948        mut f: F,
949    ) -> Result<f64, E>
950    where
951        E: From<RuntimeError>,
952        F: FnMut(Complex64) -> Result<f64, E>,
953    {
954        let mut sum = AccurateF64::zero();
955        for batch in dataset.batches() {
956            sum.push(self.try_weighted_sum_batch(params, batch, &mut f)?);
957        }
958        Ok(sum.finish())
959    }
960
961    #[cfg(test)]
962    pub(in crate::cpu) fn weighted_sum_cached<F>(
963        &self,
964        params: &ParamValues,
965        dataset: &CpuCachedDataset,
966        mut f: F,
967    ) -> RuntimeResult<f64>
968    where
969        F: FnMut(Complex64) -> f64,
970    {
971        self.try_weighted_sum_cached(params, dataset, |value| Ok(f(value)))
972    }
973
974    #[cfg(test)]
975    pub(in crate::cpu) fn try_weighted_real_sum_with_gradient_cached<E, F>(
976        &self,
977        params: &ParamValues,
978        dataset: &CpuCachedDataset,
979        mut transform: F,
980    ) -> Result<(f64, Vec<f64>), E>
981    where
982        E: From<RuntimeError>,
983        F: FnMut(Complex64) -> Result<(f64, f64), E>,
984    {
985        let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
986        for batch in dataset.batches() {
987            let (value, gradient) =
988                self.try_weighted_real_sum_with_gradient_batch(params, batch, &mut transform)?;
989            total.value.push(value);
990            for (sum, partial) in total.gradient.iter_mut().zip(gradient) {
991                sum.push(partial);
992            }
993        }
994        Ok(total.finish())
995    }
996
997    #[cfg(test)]
998    fn try_weighted_complex_sum_cached<E, F>(
999        &self,
1000        params: &ParamValues,
1001        dataset: &CpuCachedDataset,
1002        mut f: F,
1003    ) -> Result<Complex64, E>
1004    where
1005        E: From<RuntimeError>,
1006        F: FnMut(Complex64) -> Result<Complex64, E>,
1007    {
1008        let mut sum = Complex64::default();
1009        let invariant = self.scalar_invariant_values(params)?;
1010        let mut workspace = ScalarEventWorkspace::default();
1011        for batch in dataset.batches() {
1012            self.check_batch_cache(batch.cache())?;
1013            for row in 0..batch.len() {
1014                let value = self.evaluate_cache_row_prepared(
1015                    params,
1016                    batch.cache(),
1017                    row,
1018                    invariant.as_ref(),
1019                    &mut workspace,
1020                )?;
1021                sum += f(value)? * batch.weights()[row];
1022            }
1023        }
1024        Ok(sum)
1025    }
1026
1027    #[cfg(test)]
1028    pub(in crate::cpu) fn weighted_complex_sum_cached<F>(
1029        &self,
1030        params: &ParamValues,
1031        dataset: &CpuCachedDataset,
1032        mut f: F,
1033    ) -> RuntimeResult<Complex64>
1034    where
1035        F: FnMut(Complex64) -> Complex64,
1036    {
1037        self.try_weighted_complex_sum_cached(params, dataset, |value| Ok(f(value)))
1038    }
1039
1040    fn apply_reduction(&self, reduction: ReductionPlan, value: Complex64) -> RuntimeResult<f64> {
1041        if self.precision != Precision::F32 {
1042            return reduction
1043                .apply(value)
1044                .map(|output| output.value())
1045                .map_err(RuntimeError::from);
1046        }
1047        let real = value.re as f32;
1048        match reduction.transform() {
1049            ReductionTransform::Real => Ok(real as f64),
1050            ReductionTransform::PositiveReal if real > 0.0 => Ok(real as f64),
1051            ReductionTransform::LogPositiveReal if real > 0.0 => Ok(real.ln() as f64),
1052            ReductionTransform::PositiveReal | ReductionTransform::LogPositiveReal => reduction
1053                .apply(Complex64::from(real as f64))
1054                .map(|output| output.value())
1055                .map_err(RuntimeError::from),
1056        }
1057    }
1058
1059    #[cfg(test)]
1060    pub(crate) fn par_try_weighted_sum_cached<E, F>(
1061        &self,
1062        params: &ParamValues,
1063        dataset: &CpuCachedDataset,
1064        f: F,
1065    ) -> Result<f64, E>
1066    where
1067        E: From<RuntimeError> + Send,
1068        F: Fn(Complex64) -> Result<f64, E> + Send + Sync,
1069    {
1070        let mut total = AccurateF64::zero();
1071        for batch in dataset.batches() {
1072            total.push(self.par_try_weighted_sum_batch(params, batch, &f)?);
1073        }
1074        Ok(total.finish())
1075    }
1076
1077    #[cfg(test)]
1078    pub(crate) fn par_weighted_sum_cached<F>(
1079        &self,
1080        params: &ParamValues,
1081        dataset: &CpuCachedDataset,
1082        f: F,
1083    ) -> RuntimeResult<f64>
1084    where
1085        F: Fn(Complex64) -> f64 + Send + Sync,
1086    {
1087        self.par_try_weighted_sum_cached(params, dataset, |value| Ok(f(value)))
1088    }
1089
1090    #[cfg(test)]
1091    pub(in crate::cpu) fn par_try_weighted_real_sum_with_gradient_cached<E, F>(
1092        &self,
1093        params: &ParamValues,
1094        dataset: &CpuCachedDataset,
1095        transform: F,
1096    ) -> Result<(f64, Vec<f64>), E>
1097    where
1098        E: From<RuntimeError> + Send,
1099        F: Fn(Complex64) -> Result<(f64, f64), E> + Send + Sync,
1100    {
1101        let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
1102        for batch in dataset.batches() {
1103            let (value, gradient) =
1104                self.par_try_weighted_real_sum_with_gradient_batch(params, batch, &transform)?;
1105            total.value.push(value);
1106            for (sum, partial) in total.gradient.iter_mut().zip(gradient) {
1107                sum.push(partial);
1108            }
1109        }
1110        Ok(total.finish())
1111    }
1112
1113    #[cfg(test)]
1114    pub(crate) fn par_try_weighted_complex_sum_cached<E, F>(
1115        &self,
1116        params: &ParamValues,
1117        dataset: &CpuCachedDataset,
1118        f: F,
1119    ) -> Result<Complex64, E>
1120    where
1121        E: From<RuntimeError> + Send,
1122        F: Fn(Complex64) -> Result<Complex64, E> + Send + Sync,
1123    {
1124        let mut total = AccurateComplex64::zero();
1125        let invariant = self.scalar_invariant_values(params)?;
1126        for batch in dataset.batches() {
1127            self.check_batch_cache(batch.cache())?;
1128            #[cfg(feature = "jit")]
1129            let jit_cache = self
1130                .scalar_jit_kernel()
1131                .map(|_| JitScalarKernel::prepare_cache(batch.cache()));
1132            let n_blocks = batch.len().div_ceil(SCALAR_BLOCK_SIZE);
1133            let partial = (0..n_blocks)
1134                .into_par_iter()
1135                .try_fold(
1136                    || {
1137                        (
1138                            AccurateComplex64::zero(),
1139                            ScalarEventWorkspace::default(),
1140                            Vec::new(),
1141                        )
1142                    },
1143                    |(mut acc, mut workspace, mut output), block| {
1144                        let start = block * SCALAR_BLOCK_SIZE;
1145                        let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
1146                        self.evaluate_cache_block_prepared(
1147                            params,
1148                            batch.cache(),
1149                            start,
1150                            end,
1151                            invariant.as_ref(),
1152                            &mut workspace,
1153                            &mut output,
1154                            #[cfg(feature = "jit")]
1155                            jit_cache.as_ref(),
1156                        )?;
1157                        for (lane, value) in output.iter().copied().enumerate() {
1158                            acc.push(f(value)? * batch.weights()[start + lane]);
1159                        }
1160                        Ok::<_, E>((acc, workspace, output))
1161                    },
1162                )
1163                .try_reduce(
1164                    || {
1165                        (
1166                            AccurateComplex64::zero(),
1167                            ScalarEventWorkspace::default(),
1168                            Vec::new(),
1169                        )
1170                    },
1171                    |(mut lhs, workspace, output), (rhs, _, _)| {
1172                        lhs.merge(rhs);
1173                        Ok::<_, E>((lhs, workspace, output))
1174                    },
1175                )?;
1176            total.merge(partial.0);
1177        }
1178        Ok(total.finish())
1179    }
1180
1181    #[cfg(test)]
1182    pub(crate) fn par_weighted_complex_sum_cached<F>(
1183        &self,
1184        params: &ParamValues,
1185        dataset: &CpuCachedDataset,
1186        f: F,
1187    ) -> RuntimeResult<Complex64>
1188    where
1189        F: Fn(Complex64) -> Complex64 + Send + Sync,
1190    {
1191        self.par_try_weighted_complex_sum_cached(params, dataset, |value| Ok(f(value)))
1192    }
1193}