#[derive(Debug, Clone, Copy, PartialEq)]
pub enum Kernel {
Squared,
Huber(f64),
Cauchy(f64),
Tukey(f64),
GemanMcClure(f64),
}
impl Kernel {
pub fn weight(self, residual: f64) -> f64 {
let e = residual.abs();
match self {
Self::Squared => 1.0,
Self::Huber(delta) => {
if e <= delta {
1.0
} else {
delta / e
}
}
Self::Cauchy(c) => {
let t = e / c;
1.0 / (1.0 + t * t)
}
Self::Tukey(c) => {
if e <= c {
let t = e / c;
let s = 1.0 - t * t;
s * s
} else {
0.0
}
}
Self::GemanMcClure(c) => {
let denominator = c * c + e * e;
(c * c * c * c) / (denominator * denominator)
}
}
}
pub fn loss(self, residual: f64) -> f64 {
let e = residual.abs();
match self {
Self::Squared => 0.5 * e * e,
Self::Huber(delta) => {
if e <= delta {
0.5 * e * e
} else {
delta * (e - 0.5 * delta)
}
}
Self::Cauchy(c) => {
let t = e / c;
0.5 * c * c * (t * t).ln_1p()
}
Self::Tukey(c) => {
let limit = c * c / 6.0;
if e <= c {
let u = (e / c) * (e / c);
limit * u * (3.0 - 3.0 * u + u * u)
} else {
limit
}
}
Self::GemanMcClure(c) => {
let e2 = e * e;
0.5 * c * c * e2 / (c * c + e2)
}
}
}
pub fn scale(self) -> Option<f64> {
match self {
Self::Squared => None,
Self::Huber(c) | Self::Cauchy(c) | Self::Tukey(c) | Self::GemanMcClure(c) => Some(c),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const KERNELS: [Kernel; 5] = [
Kernel::Squared,
Kernel::Huber(1.0),
Kernel::Cauchy(1.0),
Kernel::Tukey(1.0),
Kernel::GemanMcClure(1.0),
];
#[test]
fn all_kernels_agree_near_zero() {
for kernel in KERNELS {
assert!((kernel.weight(0.0) - 1.0).abs() < 1e-12, "{kernel:?}");
let e = 1e-6;
assert!((kernel.weight(e) - 1.0).abs() < 1e-10, "{kernel:?}");
assert!(
(kernel.loss(e) - 0.5 * e * e).abs() < 1e-20,
"{kernel:?}: ρ = {}",
kernel.loss(e)
);
}
}
#[test]
fn weight_is_the_derivative_of_loss_over_residual() {
const H: f64 = 1e-6;
for kernel in KERNELS {
for e in [0.1, 0.5, 0.9, 1.5, 3.0, 10.0] {
let psi = (kernel.loss(e + H) - kernel.loss(e - H)) / (2.0 * H);
let expected = psi / e;
let actual = kernel.weight(e);
assert!(
(actual - expected).abs() < 1e-6,
"{kernel:?} at e = {e}: weight {actual}, but ψ/e = {expected}"
);
}
}
}
#[test]
fn robust_kernels_bound_outlier_influence() {
let far = 1e4;
assert!(Kernel::Squared.loss(far) > 1e7);
assert!(Kernel::Huber(1.0).loss(far) < 1e5);
assert!(Kernel::Cauchy(1.0).loss(far) < 20.0);
assert_eq!(Kernel::Tukey(1.0).loss(far), 1.0 / 6.0);
assert!(Kernel::GemanMcClure(1.0).loss(far) < 0.51);
assert_eq!(Kernel::Tukey(1.0).weight(1.5), 0.0);
assert!(Kernel::Huber(1.0).weight(far) < 1e-3);
}
#[test]
fn weight_is_non_increasing() {
for kernel in KERNELS {
let mut previous = f64::INFINITY;
for step in 0..200 {
let weight = kernel.weight(step as f64 * 0.05);
assert!(weight <= previous + 1e-12, "{kernel:?} at step {step}");
previous = weight;
}
}
}
}