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 mut output = vec![Complex64::ZERO; batch.len()];
307 let mut evaluate = || -> RuntimeResult<()> {
308 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 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 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 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 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}