fpm_rs/algorithms/
objective.rs1use 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}