Skip to main content

fpm_rs/algorithms/
objective.rs

1//! Optimization objectives used by reconstruction algorithms.
2
3use serde::{Deserialize, Serialize};
4
5use crate::{Result, error::Error};
6
7/// Scalar data-fidelity objective between predicted and measured intensity frames.
8#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
9pub enum LossType {
10    /// Mean squared error between predicted and measured amplitudes.
11    #[default]
12    AmplitudeMse,
13    /// Mean squared error between predicted and measured intensities.
14    IntensityMse,
15    /// Mean Poisson negative log likelihood, omitting measurement-only constants.
16    PoissonNegativeLogLikelihood,
17    /// Huber loss on amplitude residuals with the implementation's fixed unit transition.
18    HuberAmplitude,
19}
20
21/// Computes the selected mean frame loss from equal non-empty intensity slices.
22///
23/// Values must be finite and non-negative; invalid lengths or values return an error.
24pub fn loss(predicted: &[f64], measured: &[f64], loss_type: LossType) -> Result<f64> {
25    if predicted.len() != measured.len() || predicted.is_empty() {
26        return Err(Error::InvalidShape(format!(
27            "loss inputs have lengths {} and {}",
28            predicted.len(),
29            measured.len()
30        )));
31    }
32    let mut sum = 0.0;
33    for (&predicted, &measured) in predicted.iter().zip(measured) {
34        if !predicted.is_finite() || !measured.is_finite() {
35            return Err(Error::Numerical(
36                "loss input contains a non-finite value".into(),
37            ));
38        }
39        sum += point_loss(predicted, measured, loss_type);
40    }
41    Ok(sum / predicted.len() as f64)
42}
43
44pub(crate) fn point_loss(predicted: f64, measured: f64, loss_type: LossType) -> f64 {
45    match loss_type {
46        LossType::AmplitudeMse => {
47            let residual = predicted.max(0.0).sqrt() - measured.max(0.0).sqrt();
48            residual * residual
49        }
50        LossType::IntensityMse => {
51            let residual = predicted - measured;
52            residual * residual
53        }
54        LossType::PoissonNegativeLogLikelihood => {
55            let mean = predicted.max(1e-12);
56            mean - measured.max(0.0) * mean.ln()
57        }
58        LossType::HuberAmplitude => {
59            let residual = predicted.max(0.0).sqrt() - measured.max(0.0).sqrt();
60            let absolute = residual.abs();
61            if absolute <= 1.0 {
62                0.5 * residual * residual
63            } else {
64                absolute - 0.5
65            }
66        }
67    }
68}