#[derive(Clone, Debug, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum SymRegLoss {
#[default]
Mse,
Huber {
delta: f64,
},
TrimmedMse {
alpha: f64,
},
}
pub(super) fn huber_loss(residuals: &[f64], delta: f64) -> f64 {
if residuals.is_empty() {
return 0.0;
}
let sum: f64 = residuals
.iter()
.map(|&r| {
let ar = r.abs();
if ar <= delta {
0.5 * r * r
} else {
delta * (ar - 0.5 * delta)
}
})
.sum();
sum / residuals.len() as f64
}
pub(super) fn huber_grad_factor(r: f64, delta: f64) -> f64 {
if r.abs() <= delta {
r
} else {
delta * r.signum()
}
}
pub(super) fn trimmed_mse(residuals: &[f64], alpha: f64) -> f64 {
if residuals.is_empty() {
return 0.0;
}
let mut sorted: Vec<f64> = residuals.iter().map(|r| r * r).collect();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let keep = ((1.0 - alpha) * sorted.len() as f64).ceil() as usize;
let keep = keep.max(1).min(sorted.len());
sorted[..keep].iter().sum::<f64>() / keep as f64
}
pub(super) fn trimmed_mse_grad_factor(r: f64, residuals: &[f64], alpha: f64) -> f64 {
if residuals.is_empty() || alpha <= 0.0 {
return r;
}
let mut abs_res: Vec<f64> = residuals.iter().map(|x| x.abs()).collect();
abs_res.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let q_idx = ((1.0 - alpha) * (abs_res.len() - 1) as f64).round() as usize;
let q = abs_res[q_idx.min(abs_res.len() - 1)].max(1e-12);
let sharpness = 3.0_f64;
let w = 1.0 / (1.0 + (r.abs() / q - (1.0 - alpha)).exp() * sharpness.exp());
w.clamp(0.0, 1.0) * r
}