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)]
9pub enum LossType {
10 #[default]
12 AmplitudeMse,
13 IntensityMse,
15 PoissonNegativeLogLikelihood,
17 HuberAmplitude,
19}
20
21pub 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}