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
26struct 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 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 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 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 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}