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 mut output = vec![Complex64::ZERO; batch.len()];
307        let mut evaluate = || -> RuntimeResult<()> {
308            // Balance contiguous event tiles even when the pool size is not a
309            // power of two. Each tile reuses scratch across its event blocks.
310            let tile_len = block_count
311                .div_ceil(rayon::current_num_threads())
312                .saturating_mul(SCALAR_BLOCK_SIZE);
313            output
314                .par_chunks_mut(tile_len)
315                .enumerate()
316                .try_for_each_init(
317                    || {
318                        (
319                            ScalarEventWorkspace::default(),
320                            Vec::with_capacity(SCALAR_BLOCK_SIZE),
321                        )
322                    },
323                    |(workspace, block_output), (tile, target)| {
324                        for (block, target) in target.chunks_mut(SCALAR_BLOCK_SIZE).enumerate() {
325                            let start = tile * tile_len + block * SCALAR_BLOCK_SIZE;
326                            self.evaluate_cache_block_prepared(
327                                params,
328                                batch.cache(),
329                                start,
330                                start + target.len(),
331                                invariant.as_ref(),
332                                workspace,
333                                block_output,
334                                #[cfg(feature = "jit")]
335                                jit_cache.as_ref(),
336                            )?;
337                            target.copy_from_slice(block_output);
338                        }
339                        Ok(())
340                    },
341                )
342        };
343        if pool_installed {
344            evaluate()?;
345        } else {
346            execution.install(evaluate)?;
347        }
348        Ok(output)
349    }
350
351    pub(crate) fn evaluate_prepared_dataset_many(
352        &self,
353        execution: &Execution,
354        params: &[ParamValues],
355        dataset: &CpuPreparedDataset,
356    ) -> RuntimeResult<Vec<Vec<Complex64>>> {
357        self.evaluate_prepared_dataset_many_local(execution, params, dataset, None)
358            .map(|(values, _)| values)
359    }
360
361    pub(crate) fn evaluate_prepared_dataset_many_with_reduction(
362        &self,
363        execution: &Execution,
364        params: &[ParamValues],
365        dataset: &CpuPreparedDataset,
366        reduction: ReductionPlan,
367    ) -> RuntimeResult<(Vec<Vec<Complex64>>, Vec<f64>)> {
368        let local =
369            self.evaluate_prepared_dataset_many_local(execution, params, dataset, Some(reduction));
370        if !execution.all_succeeded(local.is_ok()) {
371            return local.and(Err(RuntimeError::DistributedPeerFailure));
372        }
373        let (output, sums) = local?;
374        Ok((
375            output,
376            sums.unwrap_or_default()
377                .into_iter()
378                .map(|sum| execution.sum_f64(sum))
379                .collect(),
380        ))
381    }
382
383    fn evaluate_prepared_dataset_many_local(
384        &self,
385        execution: &Execution,
386        params: &[ParamValues],
387        dataset: &CpuPreparedDataset,
388        reduction: Option<ReductionPlan>,
389    ) -> RuntimeResult<PreparedManyEvaluation> {
390        let mut stream = PreparedBatchStream::prepare(self, execution, dataset)?;
391        let mut output = params
392            .iter()
393            .map(|_| Vec::with_capacity(dataset.stats().local_events()))
394            .collect::<Vec<_>>();
395        let mut sums = reduction.map(|_| {
396            params
397                .iter()
398                .map(|_| AccurateF64::zero())
399                .collect::<Vec<_>>()
400        });
401        while let Some(batch) = stream.next()? {
402            let batch = batch.cached();
403            for (index, (parameters, values)) in params.iter().zip(&mut output).enumerate() {
404                let batch_values = self.evaluate_cache(parameters, batch.cache())?;
405                if let (Some(reduction), Some(sums)) = (reduction, sums.as_mut()) {
406                    for (weight, value) in batch.weights().iter().zip(&batch_values) {
407                        sums[index].push(*weight * reduction.apply(*value)?.value());
408                    }
409                }
410                values.extend(batch_values);
411            }
412        }
413        Ok((
414            output,
415            sums.map(|sums| sums.into_iter().map(AccurateF64::finish).collect()),
416        ))
417    }
418
419    /// Execute a weighted reduction over a prepared dataset.
420    ///
421    /// # Errors
422    ///
423    /// Returns [`RuntimeError`] when streaming, cache validation, evaluation,
424    /// or reduction fails, or another distributed worker reports failure.
425    pub fn reduce(
426        &self,
427        execution: &Execution,
428        params: &ParamValues,
429        dataset: &CpuPreparedDataset,
430        reduction: ReductionPlan,
431    ) -> RuntimeResult<f64> {
432        let local = (|| {
433            let mut stream = PreparedBatchStream::prepare(self, execution, dataset)?;
434            let mut reducer = ValueReducer::new();
435            while let Some(batch) = stream.next()? {
436                reducer.consume(self, execution, params, batch.cached(), reduction)?;
437            }
438            Ok(reducer.finish())
439        })();
440        if !execution.all_succeeded(local.is_ok()) {
441            return local.and(Err(RuntimeError::DistributedPeerFailure));
442        }
443        Ok(execution.sum_f64(local?))
444    }
445
446    /// Execute a weighted reduction and its free-parameter gradient.
447    ///
448    /// # Errors
449    ///
450    /// Returns [`RuntimeError`] when streaming, cache validation,
451    /// differentiation, evaluation, or reduction fails, or another distributed
452    /// worker reports failure.
453    pub fn reduce_with_gradient(
454        &self,
455        execution: &Execution,
456        params: &ParamValues,
457        dataset: &CpuPreparedDataset,
458        reduction: ReductionPlan,
459    ) -> RuntimeResult<ReductionEvaluation> {
460        let local = (|| {
461            let mut stream = PreparedBatchStream::prepare(self, execution, dataset)?;
462            let mut reducer = GradientReducer::new(self.free_parameter_count());
463            let transform = |value| {
464                reduction
465                    .apply(value)
466                    .map(|output| output.into_parts())
467                    .map_err(RuntimeError::from)
468            };
469            while let Some(batch) = stream.next()? {
470                reducer.consume(self, execution, params, batch.cached(), &transform)?;
471            }
472            Ok(reducer.finish())
473        })();
474        if !execution.all_succeeded(local.is_ok()) {
475            return local.and(Err(RuntimeError::DistributedPeerFailure));
476        }
477        let (value, gradient) = local?;
478        let value = execution.sum_f64(value);
479        let gradient = execution.sum_slice(&gradient);
480        Ok(ReductionEvaluation { value, gradient })
481    }
482
483    /// Evaluates every event in a fully cached dataset.
484    ///
485    /// # Errors
486    ///
487    /// Returns [`RuntimeError`] when parameters or a cache layout are
488    /// incompatible, evaluation fails, or a matrix is singular.
489    pub fn evaluate_cached_dataset(
490        &self,
491        params: &ParamValues,
492        dataset: &CpuCachedDataset,
493    ) -> RuntimeResult<Vec<Complex64>> {
494        let total_len = dataset.batches.iter().map(CpuCachedBatch::len).sum();
495        let mut out = Vec::with_capacity(total_len);
496        let invariant = self.scalar_invariant_values(params)?;
497        let mut workspace = ScalarEventWorkspace::default();
498        for batch in &dataset.batches {
499            self.check_batch_cache(batch.cache())?;
500            for row in 0..batch.len() {
501                out.push(self.evaluate_cache_row_prepared(
502                    params,
503                    batch.cache(),
504                    row,
505                    invariant.as_ref(),
506                    &mut workspace,
507                )?);
508            }
509        }
510        Ok(out)
511    }
512
513    /// Evaluates every event and gradient in a fully cached dataset.
514    ///
515    /// # Errors
516    ///
517    /// Returns [`RuntimeError`] when parameters or a cache layout are
518    /// incompatible, or differentiation or evaluation fails.
519    pub fn evaluate_cached_dataset_with_gradient(
520        &self,
521        params: &ParamValues,
522        dataset: &CpuCachedDataset,
523    ) -> RuntimeResult<Vec<ValueGradient>> {
524        let total_len = dataset.batches.iter().map(CpuCachedBatch::len).sum();
525        let mut out = Vec::with_capacity(total_len);
526        for batch in &dataset.batches {
527            out.extend(self.evaluate_cache_with_gradient(params, batch.cache())?);
528        }
529        Ok(out)
530    }
531
532    fn try_weighted_sum_batch<E, F>(
533        &self,
534        params: &ParamValues,
535        batch: &CpuCachedBatch,
536        mut f: F,
537    ) -> Result<f64, E>
538    where
539        E: From<RuntimeError>,
540        F: FnMut(Complex64) -> Result<f64, E>,
541    {
542        self.check_batch_cache(batch.cache())?;
543        let invariant = self.scalar_invariant_values(params)?;
544        let mut workspace = ScalarEventWorkspace::default();
545        let mut output = Vec::with_capacity(SCALAR_BLOCK_SIZE);
546        #[cfg(feature = "jit")]
547        let jit_cache = self
548            .scalar_jit_kernel()
549            .map(|_| JitScalarKernel::prepare_cache(batch.cache()));
550        let mut sum = AccurateF64::zero();
551        for start in (0..batch.len()).step_by(SCALAR_BLOCK_SIZE) {
552            let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
553            self.evaluate_cache_block_prepared(
554                params,
555                batch.cache(),
556                start,
557                end,
558                invariant.as_ref(),
559                &mut workspace,
560                &mut output,
561                #[cfg(feature = "jit")]
562                jit_cache.as_ref(),
563            )?;
564            for (lane, value) in output.iter().copied().enumerate() {
565                sum.push(batch.weights()[start + lane] * f(value)?);
566            }
567        }
568        Ok(sum.finish())
569    }
570
571    fn par_try_weighted_sum_batch<E, F>(
572        &self,
573        params: &ParamValues,
574        batch: &CpuCachedBatch,
575        f: F,
576    ) -> Result<f64, E>
577    where
578        E: From<RuntimeError> + Send,
579        F: Fn(Complex64) -> Result<f64, E> + Send + Sync,
580    {
581        self.check_batch_cache(batch.cache())?;
582        let invariant = self.scalar_invariant_values(params)?;
583        #[cfg(feature = "jit")]
584        let jit_cache = self
585            .scalar_jit_kernel()
586            .map(|_| JitScalarKernel::prepare_cache(batch.cache()));
587        let n_blocks = batch.len().div_ceil(SCALAR_BLOCK_SIZE);
588        let total = (0..n_blocks)
589            .into_par_iter()
590            .try_fold(
591                || {
592                    (
593                        AccurateF64::zero(),
594                        ScalarEventWorkspace::default(),
595                        Vec::new(),
596                    )
597                },
598                |(mut acc, mut workspace, mut output), block| {
599                    let start = block * SCALAR_BLOCK_SIZE;
600                    let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
601                    self.evaluate_cache_block_prepared(
602                        params,
603                        batch.cache(),
604                        start,
605                        end,
606                        invariant.as_ref(),
607                        &mut workspace,
608                        &mut output,
609                        #[cfg(feature = "jit")]
610                        jit_cache.as_ref(),
611                    )?;
612                    for (lane, value) in output.iter().copied().enumerate() {
613                        acc.push(batch.weights()[start + lane] * f(value)?);
614                    }
615                    Ok::<_, E>((acc, workspace, output))
616                },
617            )
618            .try_reduce(
619                || {
620                    (
621                        AccurateF64::zero(),
622                        ScalarEventWorkspace::default(),
623                        Vec::new(),
624                    )
625                },
626                |(mut lhs, workspace, output), (rhs, _, _)| {
627                    lhs.merge(rhs);
628                    Ok::<_, E>((lhs, workspace, output))
629                },
630            )?;
631        Ok(total.0.finish())
632    }
633
634    fn try_weighted_real_sum_with_gradient_batch<E, F>(
635        &self,
636        params: &ParamValues,
637        batch: &CpuCachedBatch,
638        transform: &mut F,
639    ) -> Result<(f64, Vec<f64>), E>
640    where
641        E: From<RuntimeError>,
642        F: FnMut(Complex64) -> Result<(f64, f64), E>,
643    {
644        self.check_batch_cache(batch.cache())?;
645        #[cfg(feature = "jit")]
646        if let (Some(value_kernel), Some(gradient_kernel)) =
647            (self.scalar_jit_kernel(), self.gradient_jit_kernel())
648        {
649            return self.try_weighted_real_sum_with_jit_gradient_batch(
650                params,
651                batch,
652                transform,
653                value_kernel,
654                gradient_kernel,
655            );
656        }
657        if self.precision != Precision::F32
658            && let Some(interpreter) = self.gradient_interpreter()
659            && let Some(mut state) = interpreter.prepare_real_blocks(params)?
660        {
661            let output_count = state.output_count();
662            let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
663            for block in 0..batch.len().div_ceil(SCALAR_BLOCK_SIZE) {
664                let start = block * SCALAR_BLOCK_SIZE;
665                let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
666                let outputs = state.evaluate(batch.cache(), start, end)?;
667                for (lane, row) in outputs.chunks_exact(output_count).enumerate() {
668                    let (value, derivative) = transform(row[0])?;
669                    total.push(batch.weights()[start + lane], value, derivative, &row[1..]);
670                }
671            }
672            return Ok(total.finish());
673        }
674        if self.precision == Precision::F32
675            && let Some(ir) = self.f32_gradient_fallback_real.as_ref()
676        {
677            let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
678            let mut gradient = Vec::new();
679            for row in 0..batch.len() {
680                let (value, model_gradient) = self.evaluate_f32_gradient_component_prepared(
681                    ir,
682                    params,
683                    F32KernelInput::Cache(Some((batch.cache(), row))),
684                    &mut gradient,
685                )?;
686                let (value, derivative) = transform(value)?;
687                total.push_f32(batch.weights()[row], value, derivative, model_gradient);
688            }
689            return Ok(total.finish());
690        }
691        let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
692        for row in 0..batch.len() {
693            let evaluation =
694                self.evaluate_cache_row_with_gradient_unchecked(params, batch.cache(), row)?;
695            let (value, derivative) = transform(evaluation.value())?;
696            total.push(
697                batch.weights()[row],
698                value,
699                derivative,
700                evaluation.gradient(),
701            );
702        }
703        Ok(total.finish())
704    }
705
706    fn par_try_weighted_real_sum_with_gradient_batch<E, F>(
707        &self,
708        params: &ParamValues,
709        batch: &CpuCachedBatch,
710        transform: &F,
711    ) -> Result<(f64, Vec<f64>), E>
712    where
713        E: From<RuntimeError> + Send,
714        F: Fn(Complex64) -> Result<(f64, f64), E> + Send + Sync,
715    {
716        #[cfg(feature = "jit")]
717        if let (Some(value_kernel), Some(gradient_kernel)) =
718            (self.scalar_jit_kernel(), self.gradient_jit_kernel())
719        {
720            return self.par_try_weighted_real_sum_with_jit_gradient_batch(
721                params,
722                batch,
723                transform,
724                value_kernel,
725                gradient_kernel,
726            );
727        }
728        if self.precision != Precision::F32
729            && let Some(interpreter) = self.gradient_interpreter()
730            && let Some(state) = interpreter.prepare_real_blocks(params)?
731        {
732            let output_count = state.output_count();
733            let partial = (0..batch.len().div_ceil(SCALAR_BLOCK_SIZE))
734                .into_par_iter()
735                .try_fold(
736                    || {
737                        (
738                            RealGradientAccumulator::zero(self.free_parameter_count()),
739                            state.clone(),
740                        )
741                    },
742                    |(mut accumulator, mut state), block| {
743                        let start = block * SCALAR_BLOCK_SIZE;
744                        let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
745                        let outputs = state.evaluate(batch.cache(), start, end)?;
746                        for (lane, row) in outputs.chunks_exact(output_count).enumerate() {
747                            let (value, derivative) = transform(row[0])?;
748                            accumulator.push(
749                                batch.weights()[start + lane],
750                                value,
751                                derivative,
752                                &row[1..],
753                            );
754                        }
755                        Ok::<_, E>((accumulator, state))
756                    },
757                )
758                .try_reduce(
759                    || {
760                        (
761                            RealGradientAccumulator::zero(self.free_parameter_count()),
762                            state.clone(),
763                        )
764                    },
765                    |(mut lhs, state), (rhs, _)| {
766                        lhs.merge(rhs);
767                        Ok::<_, E>((lhs, state))
768                    },
769                )?;
770            return Ok(partial.0.finish());
771        }
772        if self.precision == Precision::F32
773            && let Some(ir) = self.f32_gradient_fallback_real.as_ref()
774        {
775            let partial = (0..batch.len())
776                .into_par_iter()
777                .try_fold(
778                    || {
779                        (
780                            RealGradientAccumulator::zero(self.free_parameter_count()),
781                            Vec::new(),
782                        )
783                    },
784                    |(mut accumulator, mut gradient), row| {
785                        let (value, model_gradient) = self
786                            .evaluate_f32_gradient_component_prepared(
787                                ir,
788                                params,
789                                F32KernelInput::Cache(Some((batch.cache(), row))),
790                                &mut gradient,
791                            )?;
792                        let (value, derivative) = transform(value)?;
793                        accumulator.push_f32(
794                            batch.weights()[row],
795                            value,
796                            derivative,
797                            model_gradient,
798                        );
799                        Ok::<_, E>((accumulator, gradient))
800                    },
801                )
802                .try_reduce(
803                    || {
804                        (
805                            RealGradientAccumulator::zero(self.free_parameter_count()),
806                            Vec::new(),
807                        )
808                    },
809                    |(mut lhs, gradient), (rhs, _)| {
810                        lhs.merge(rhs);
811                        Ok::<_, E>((lhs, gradient))
812                    },
813                )?;
814            return Ok(partial.0.finish());
815        }
816        let partial = (0..batch.len())
817            .into_par_iter()
818            .try_fold(
819                || RealGradientAccumulator::zero(self.free_parameter_count()),
820                |mut accumulator, row| {
821                    let evaluation = self.evaluate_cache_row_with_gradient_unchecked(
822                        params,
823                        batch.cache(),
824                        row,
825                    )?;
826                    let (value, derivative) = transform(evaluation.value())?;
827                    accumulator.push(
828                        batch.weights()[row],
829                        value,
830                        derivative,
831                        evaluation.gradient(),
832                    );
833                    Ok::<_, E>(accumulator)
834                },
835            )
836            .try_reduce(
837                || RealGradientAccumulator::zero(self.free_parameter_count()),
838                |mut lhs, rhs| {
839                    lhs.merge(rhs);
840                    Ok::<_, E>(lhs)
841                },
842            )?;
843        Ok(partial.finish())
844    }
845
846    #[cfg(feature = "jit")]
847    fn try_weighted_real_sum_with_jit_gradient_batch<E, F>(
848        &self,
849        params: &ParamValues,
850        batch: &CpuCachedBatch,
851        transform: &mut F,
852        value_kernel: &JitScalarKernel,
853        gradient_kernel: &JitGradientKernel,
854    ) -> Result<(f64, Vec<f64>), E>
855    where
856        E: From<RuntimeError>,
857        F: FnMut(Complex64) -> Result<(f64, f64), E>,
858    {
859        let cache = JitScalarKernel::prepare_cache(batch.cache());
860        let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
861        let mut values = Vec::new();
862        let mut tangents = Vec::new();
863        let mut derivatives = Vec::new();
864        for block in 0..batch.len().div_ceil(SCALAR_BLOCK_SIZE) {
865            let start = block * SCALAR_BLOCK_SIZE;
866            let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
867            value_kernel.evaluate_prepared(params, &cache, start, end, &mut values)?;
868            derivatives.clear();
869            derivatives.reserve(values.len());
870            for (lane, value) in values.iter().copied().enumerate() {
871                let (value, derivative) = transform(value)?;
872                total.value.push(batch.weights()[start + lane] * value);
873                derivatives.push(batch.weights()[start + lane] * derivative);
874            }
875            gradient_kernel.evaluate_prepared(params, &cache, start, end, 0, &mut tangents)?;
876            for (lane, factor) in derivatives.iter().enumerate() {
877                for free_index in 0..self.free_parameter_count() {
878                    total.gradient[free_index]
879                        .push(factor * tangents[lane * self.free_parameter_count() + free_index]);
880                }
881            }
882        }
883        Ok(total.finish())
884    }
885
886    #[cfg(feature = "jit")]
887    fn par_try_weighted_real_sum_with_jit_gradient_batch<E, F>(
888        &self,
889        params: &ParamValues,
890        batch: &CpuCachedBatch,
891        transform: &F,
892        value_kernel: &JitScalarKernel,
893        gradient_kernel: &JitGradientKernel,
894    ) -> Result<(f64, Vec<f64>), E>
895    where
896        E: From<RuntimeError> + Send,
897        F: Fn(Complex64) -> Result<(f64, f64), E> + Send + Sync,
898    {
899        let cache = JitScalarKernel::prepare_cache(batch.cache());
900        let partial = (0..batch.len().div_ceil(SCALAR_BLOCK_SIZE))
901            .into_par_iter()
902            .try_fold(
903                || {
904                    (
905                        RealGradientAccumulator::zero(self.free_parameter_count()),
906                        Vec::new(),
907                        Vec::new(),
908                        Vec::new(),
909                    )
910                },
911                |(mut accumulator, mut values, mut tangents, mut derivatives), block| {
912                    let start = block * SCALAR_BLOCK_SIZE;
913                    let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
914                    value_kernel.evaluate_prepared(params, &cache, start, end, &mut values)?;
915                    derivatives.clear();
916                    for (lane, value) in values.iter().copied().enumerate() {
917                        let (value, derivative) = transform(value)?;
918                        let weight = batch.weights()[start + lane];
919                        accumulator.value.push(weight * value);
920                        derivatives.push(weight * derivative);
921                    }
922                    gradient_kernel.evaluate_prepared(
923                        params,
924                        &cache,
925                        start,
926                        end,
927                        0,
928                        &mut tangents,
929                    )?;
930                    for (lane, factor) in derivatives.iter().enumerate() {
931                        for free_index in 0..self.free_parameter_count() {
932                            accumulator.gradient[free_index].push(
933                                factor * tangents[lane * self.free_parameter_count() + free_index],
934                            );
935                        }
936                    }
937                    Ok::<_, E>((accumulator, values, tangents, derivatives))
938                },
939            )
940            .try_reduce(
941                || {
942                    (
943                        RealGradientAccumulator::zero(self.free_parameter_count()),
944                        Vec::new(),
945                        Vec::new(),
946                        Vec::new(),
947                    )
948                },
949                |(mut lhs, values, tangents, derivatives), (rhs, _, _, _)| {
950                    lhs.merge(rhs);
951                    Ok::<_, E>((lhs, values, tangents, derivatives))
952                },
953            )?;
954        Ok(partial.0.finish())
955    }
956
957    #[cfg(test)]
958    fn try_weighted_sum_cached<E, F>(
959        &self,
960        params: &ParamValues,
961        dataset: &CpuCachedDataset,
962        mut f: F,
963    ) -> Result<f64, E>
964    where
965        E: From<RuntimeError>,
966        F: FnMut(Complex64) -> Result<f64, E>,
967    {
968        let mut sum = AccurateF64::zero();
969        for batch in dataset.batches() {
970            sum.push(self.try_weighted_sum_batch(params, batch, &mut f)?);
971        }
972        Ok(sum.finish())
973    }
974
975    #[cfg(test)]
976    pub(in crate::cpu) fn weighted_sum_cached<F>(
977        &self,
978        params: &ParamValues,
979        dataset: &CpuCachedDataset,
980        mut f: F,
981    ) -> RuntimeResult<f64>
982    where
983        F: FnMut(Complex64) -> f64,
984    {
985        self.try_weighted_sum_cached(params, dataset, |value| Ok(f(value)))
986    }
987
988    #[cfg(test)]
989    pub(in crate::cpu) fn try_weighted_real_sum_with_gradient_cached<E, F>(
990        &self,
991        params: &ParamValues,
992        dataset: &CpuCachedDataset,
993        mut transform: F,
994    ) -> Result<(f64, Vec<f64>), E>
995    where
996        E: From<RuntimeError>,
997        F: FnMut(Complex64) -> Result<(f64, f64), E>,
998    {
999        let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
1000        for batch in dataset.batches() {
1001            let (value, gradient) =
1002                self.try_weighted_real_sum_with_gradient_batch(params, batch, &mut transform)?;
1003            total.value.push(value);
1004            for (sum, partial) in total.gradient.iter_mut().zip(gradient) {
1005                sum.push(partial);
1006            }
1007        }
1008        Ok(total.finish())
1009    }
1010
1011    #[cfg(test)]
1012    fn try_weighted_complex_sum_cached<E, F>(
1013        &self,
1014        params: &ParamValues,
1015        dataset: &CpuCachedDataset,
1016        mut f: F,
1017    ) -> Result<Complex64, E>
1018    where
1019        E: From<RuntimeError>,
1020        F: FnMut(Complex64) -> Result<Complex64, E>,
1021    {
1022        let mut sum = Complex64::default();
1023        let invariant = self.scalar_invariant_values(params)?;
1024        let mut workspace = ScalarEventWorkspace::default();
1025        for batch in dataset.batches() {
1026            self.check_batch_cache(batch.cache())?;
1027            for row in 0..batch.len() {
1028                let value = self.evaluate_cache_row_prepared(
1029                    params,
1030                    batch.cache(),
1031                    row,
1032                    invariant.as_ref(),
1033                    &mut workspace,
1034                )?;
1035                sum += f(value)? * batch.weights()[row];
1036            }
1037        }
1038        Ok(sum)
1039    }
1040
1041    #[cfg(test)]
1042    pub(in crate::cpu) fn weighted_complex_sum_cached<F>(
1043        &self,
1044        params: &ParamValues,
1045        dataset: &CpuCachedDataset,
1046        mut f: F,
1047    ) -> RuntimeResult<Complex64>
1048    where
1049        F: FnMut(Complex64) -> Complex64,
1050    {
1051        self.try_weighted_complex_sum_cached(params, dataset, |value| Ok(f(value)))
1052    }
1053
1054    fn apply_reduction(&self, reduction: ReductionPlan, value: Complex64) -> RuntimeResult<f64> {
1055        if self.precision != Precision::F32 {
1056            return reduction
1057                .apply(value)
1058                .map(|output| output.value())
1059                .map_err(RuntimeError::from);
1060        }
1061        let real = value.re as f32;
1062        match reduction.transform() {
1063            ReductionTransform::Real => Ok(real as f64),
1064            ReductionTransform::PositiveReal if real > 0.0 => Ok(real as f64),
1065            ReductionTransform::LogPositiveReal if real > 0.0 => Ok(real.ln() as f64),
1066            ReductionTransform::PositiveReal | ReductionTransform::LogPositiveReal => reduction
1067                .apply(Complex64::from(real as f64))
1068                .map(|output| output.value())
1069                .map_err(RuntimeError::from),
1070        }
1071    }
1072
1073    #[cfg(test)]
1074    pub(crate) fn par_try_weighted_sum_cached<E, F>(
1075        &self,
1076        params: &ParamValues,
1077        dataset: &CpuCachedDataset,
1078        f: F,
1079    ) -> Result<f64, E>
1080    where
1081        E: From<RuntimeError> + Send,
1082        F: Fn(Complex64) -> Result<f64, E> + Send + Sync,
1083    {
1084        let mut total = AccurateF64::zero();
1085        for batch in dataset.batches() {
1086            total.push(self.par_try_weighted_sum_batch(params, batch, &f)?);
1087        }
1088        Ok(total.finish())
1089    }
1090
1091    #[cfg(test)]
1092    pub(crate) fn par_weighted_sum_cached<F>(
1093        &self,
1094        params: &ParamValues,
1095        dataset: &CpuCachedDataset,
1096        f: F,
1097    ) -> RuntimeResult<f64>
1098    where
1099        F: Fn(Complex64) -> f64 + Send + Sync,
1100    {
1101        self.par_try_weighted_sum_cached(params, dataset, |value| Ok(f(value)))
1102    }
1103
1104    #[cfg(test)]
1105    pub(in crate::cpu) fn par_try_weighted_real_sum_with_gradient_cached<E, F>(
1106        &self,
1107        params: &ParamValues,
1108        dataset: &CpuCachedDataset,
1109        transform: F,
1110    ) -> Result<(f64, Vec<f64>), E>
1111    where
1112        E: From<RuntimeError> + Send,
1113        F: Fn(Complex64) -> Result<(f64, f64), E> + Send + Sync,
1114    {
1115        let mut total = RealGradientAccumulator::zero(self.free_parameter_count());
1116        for batch in dataset.batches() {
1117            let (value, gradient) =
1118                self.par_try_weighted_real_sum_with_gradient_batch(params, batch, &transform)?;
1119            total.value.push(value);
1120            for (sum, partial) in total.gradient.iter_mut().zip(gradient) {
1121                sum.push(partial);
1122            }
1123        }
1124        Ok(total.finish())
1125    }
1126
1127    #[cfg(test)]
1128    pub(crate) fn par_try_weighted_complex_sum_cached<E, F>(
1129        &self,
1130        params: &ParamValues,
1131        dataset: &CpuCachedDataset,
1132        f: F,
1133    ) -> Result<Complex64, E>
1134    where
1135        E: From<RuntimeError> + Send,
1136        F: Fn(Complex64) -> Result<Complex64, E> + Send + Sync,
1137    {
1138        let mut total = AccurateComplex64::zero();
1139        let invariant = self.scalar_invariant_values(params)?;
1140        for batch in dataset.batches() {
1141            self.check_batch_cache(batch.cache())?;
1142            #[cfg(feature = "jit")]
1143            let jit_cache = self
1144                .scalar_jit_kernel()
1145                .map(|_| JitScalarKernel::prepare_cache(batch.cache()));
1146            let n_blocks = batch.len().div_ceil(SCALAR_BLOCK_SIZE);
1147            let partial = (0..n_blocks)
1148                .into_par_iter()
1149                .try_fold(
1150                    || {
1151                        (
1152                            AccurateComplex64::zero(),
1153                            ScalarEventWorkspace::default(),
1154                            Vec::new(),
1155                        )
1156                    },
1157                    |(mut acc, mut workspace, mut output), block| {
1158                        let start = block * SCALAR_BLOCK_SIZE;
1159                        let end = (start + SCALAR_BLOCK_SIZE).min(batch.len());
1160                        self.evaluate_cache_block_prepared(
1161                            params,
1162                            batch.cache(),
1163                            start,
1164                            end,
1165                            invariant.as_ref(),
1166                            &mut workspace,
1167                            &mut output,
1168                            #[cfg(feature = "jit")]
1169                            jit_cache.as_ref(),
1170                        )?;
1171                        for (lane, value) in output.iter().copied().enumerate() {
1172                            acc.push(f(value)? * batch.weights()[start + lane]);
1173                        }
1174                        Ok::<_, E>((acc, workspace, output))
1175                    },
1176                )
1177                .try_reduce(
1178                    || {
1179                        (
1180                            AccurateComplex64::zero(),
1181                            ScalarEventWorkspace::default(),
1182                            Vec::new(),
1183                        )
1184                    },
1185                    |(mut lhs, workspace, output), (rhs, _, _)| {
1186                        lhs.merge(rhs);
1187                        Ok::<_, E>((lhs, workspace, output))
1188                    },
1189                )?;
1190            total.merge(partial.0);
1191        }
1192        Ok(total.finish())
1193    }
1194
1195    #[cfg(test)]
1196    pub(crate) fn par_weighted_complex_sum_cached<F>(
1197        &self,
1198        params: &ParamValues,
1199        dataset: &CpuCachedDataset,
1200        f: F,
1201    ) -> RuntimeResult<Complex64>
1202    where
1203        F: Fn(Complex64) -> Complex64 + Send + Sync,
1204    {
1205        self.par_try_weighted_complex_sum_cached(params, dataset, |value| Ok(f(value)))
1206    }
1207}