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#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
8pub enum LossType {
9    #[default]
10    AmplitudeMse,
11    IntensityMse,
12    PoissonNegativeLogLikelihood,
13    HuberAmplitude,
14}
15
16pub fn loss(predicted: &[f64], measured: &[f64], loss_type: LossType) -> Result<f64> {
17    if predicted.len() != measured.len() || predicted.is_empty() {
18        return Err(Error::InvalidShape(format!(
19            "loss inputs have lengths {} and {}",
20            predicted.len(),
21            measured.len()
22        )));
23    }
24    let mut sum = 0.0;
25    for (&predicted, &measured) in predicted.iter().zip(measured) {
26        if !predicted.is_finite() || !measured.is_finite() {
27            return Err(Error::Numerical(
28                "loss input contains a non-finite value".into(),
29            ));
30        }
31        sum += point_loss(predicted, measured, loss_type);
32    }
33    Ok(sum / predicted.len() as f64)
34}
35
36pub(crate) fn point_loss(predicted: f64, measured: f64, loss_type: LossType) -> f64 {
37    match loss_type {
38        LossType::AmplitudeMse => {
39            let residual = predicted.max(0.0).sqrt() - measured.max(0.0).sqrt();
40            residual * residual
41        }
42        LossType::IntensityMse => {
43            let residual = predicted - measured;
44            residual * residual
45        }
46        LossType::PoissonNegativeLogLikelihood => {
47            let mean = predicted.max(1e-12);
48            mean - measured.max(0.0) * mean.ln()
49        }
50        LossType::HuberAmplitude => {
51            let residual = predicted.max(0.0).sqrt() - measured.max(0.0).sqrt();
52            let absolute = residual.abs();
53            if absolute <= 1.0 {
54                0.5 * residual * residual
55            } else {
56                absolute - 0.5
57            }
58        }
59    }
60}