use super::{
column_quantiles, fitted, for_each_row, handle_zero_scale, validate_matrix,
validate_transform_matrix,
};
use crate::error::Error;
use crate::{Deserialize, Serialize};
use ndarray::{Array1, Array2, ArrayBase, Data, Ix2};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RobustScaler {
with_centering: bool,
with_scaling: bool,
quantile_range: (f64, f64),
center: Option<Array1<f64>>,
scale: Option<Array1<f64>>,
n_samples_seen: usize,
}
impl Default for RobustScaler {
fn default() -> Self {
Self::new()
}
}
impl RobustScaler {
pub fn new() -> Self {
Self {
with_centering: true,
with_scaling: true,
quantile_range: (25.0, 75.0),
center: None,
scale: None,
n_samples_seen: 0,
}
}
pub fn with_centering(mut self, with_centering: bool) -> Self {
self.with_centering = with_centering;
self
}
pub fn with_scaling(mut self, with_scaling: bool) -> Self {
self.with_scaling = with_scaling;
self
}
pub fn with_quantile_range(mut self, low: f64, high: f64) -> Result<Self, Error> {
if !low.is_finite() || !high.is_finite() || low < 0.0 || high > 100.0 || low >= high {
return Err(Error::invalid_parameter(
"quantile_range",
format!(
"must satisfy 0 <= low < high <= 100, got ({}, {})",
low, high
),
));
}
self.quantile_range = (low, high);
self.center = None;
self.scale = None;
self.n_samples_seen = 0;
Ok(self)
}
get_field!(get_with_centering, with_centering, bool);
get_field!(get_with_scaling, with_scaling, bool);
get_field!(get_quantile_range, quantile_range, (f64, f64));
get_field!(get_n_samples_seen, n_samples_seen, usize);
get_field_as_ref!(get_center, center, Option<&Array1<f64>>);
get_field_as_ref!(get_scale, scale, Option<&Array1<f64>>);
#[inline]
pub fn get_n_features(&self) -> Option<usize> {
self.center.as_ref().map(|center| center.len())
}
pub fn fit<S>(&mut self, x: &ArrayBase<S, Ix2>) -> Result<&mut Self, Error>
where
S: Data<Elem = f64>,
{
validate_matrix(x)?;
let (low, high) = self.quantile_range;
let per_feature = column_quantiles(&x.view(), &[low / 100.0, 0.5, high / 100.0]);
self.center = Some(per_feature.iter().map(|q| q[1]).collect());
self.scale = Some(
per_feature
.iter()
.map(|q| handle_zero_scale(q[2] - q[0]))
.collect(),
);
self.n_samples_seen = x.nrows();
Ok(self)
}
pub fn transform<S>(&self, x: &ArrayBase<S, Ix2>) -> Result<Array2<f64>, Error>
where
S: Data<Elem = f64>,
{
let (center, scale) = self.fitted_stats()?;
validate_transform_matrix(x, center.len())?;
let mut result = x.to_owned();
let (with_centering, with_scaling) = (self.with_centering, self.with_scaling);
for_each_row(&mut result, |mut row| {
match (with_centering, with_scaling) {
(true, true) => {
for ((value, &c), &s) in row.iter_mut().zip(center).zip(scale) {
*value = (*value - c) / s;
}
}
(true, false) => {
for (value, &c) in row.iter_mut().zip(center) {
*value -= c;
}
}
(false, true) => {
for (value, &s) in row.iter_mut().zip(scale) {
*value /= s;
}
}
(false, false) => {}
}
});
Ok(result)
}
pub fn fit_transform<S>(&mut self, x: &ArrayBase<S, Ix2>) -> Result<Array2<f64>, Error>
where
S: Data<Elem = f64>,
{
self.fit(x)?;
self.transform(x)
}
pub fn inverse_transform<S>(&self, x: &ArrayBase<S, Ix2>) -> Result<Array2<f64>, Error>
where
S: Data<Elem = f64>,
{
let (center, scale) = self.fitted_stats()?;
validate_transform_matrix(x, center.len())?;
let mut result = x.to_owned();
let (with_centering, with_scaling) = (self.with_centering, self.with_scaling);
for_each_row(&mut result, |mut row| {
match (with_centering, with_scaling) {
(true, true) => {
for ((value, &c), &s) in row.iter_mut().zip(center).zip(scale) {
*value = *value * s + c;
}
}
(true, false) => {
for (value, &c) in row.iter_mut().zip(center) {
*value += c;
}
}
(false, true) => {
for (value, &s) in row.iter_mut().zip(scale) {
*value *= s;
}
}
(false, false) => {}
}
});
Ok(result)
}
model_save_and_load_methods!(RobustScaler);
fn fitted_stats(&self) -> Result<(&Array1<f64>, &Array1<f64>), Error> {
Ok((
fitted(&self.center, "RobustScaler")?,
fitted(&self.scale, "RobustScaler")?,
))
}
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
#[test]
fn matches_scikit_learn_reference_example() {
let x = array![[1.0, -2.0, 2.0], [-2.0, 1.0, 3.0], [4.0, 1.0, -2.0]];
let z = RobustScaler::new().fit_transform(&x).unwrap();
let expected = array![[0.0, -2.0, 0.0], [-1.0, 0.0, 0.4], [1.0, 0.0, -1.6]];
for (actual, want) in z.iter().zip(expected.iter()) {
assert!((actual - want).abs() < 1e-12, "{actual} != {want}");
}
}
#[test]
fn learns_median_and_iqr() {
let x = array![[1.0], [2.0], [3.0], [4.0]];
let mut scaler = RobustScaler::new();
scaler.fit(&x).unwrap();
assert!((scaler.get_center().unwrap()[0] - 2.5).abs() < 1e-12);
assert!((scaler.get_scale().unwrap()[0] - 1.5).abs() < 1e-12);
assert_eq!(scaler.get_n_samples_seen(), 4);
assert_eq!(scaler.get_n_features(), Some(1));
}
#[test]
fn resists_an_outlier() {
let clean = array![
[1.0],
[2.0],
[3.0],
[4.0],
[5.0],
[6.0],
[7.0],
[8.0],
[9.0]
];
let spoiled = array![
[1.0],
[2.0],
[3.0],
[4.0],
[5.0],
[6.0],
[7.0],
[8.0],
[1000.0]
];
let mut on_clean = RobustScaler::new();
on_clean.fit(&clean).unwrap();
let mut on_spoiled = RobustScaler::new();
on_spoiled.fit(&spoiled).unwrap();
assert_eq!(on_clean.get_center(), on_spoiled.get_center());
assert_eq!(on_clean.get_scale(), on_spoiled.get_scale());
assert!((on_spoiled.get_center().unwrap()[0] - 5.0).abs() < 1e-12);
assert!((on_spoiled.get_scale().unwrap()[0] - 4.0).abs() < 1e-12);
let mean = spoiled.iter().sum::<f64>() / spoiled.len() as f64;
assert!(mean > 100.0);
}
#[test]
fn custom_quantile_range() {
let x = array![[1.0], [2.0], [3.0], [4.0], [5.0]];
let mut narrow = RobustScaler::new();
narrow.fit(&x).unwrap();
let iqr = narrow.get_scale().unwrap()[0];
let mut wide = RobustScaler::new().with_quantile_range(10.0, 90.0).unwrap();
assert!(wide.get_center().is_none(), "changing the range unfits");
wide.fit(&x).unwrap();
assert!(wide.get_scale().unwrap()[0] > iqr);
assert_eq!(wide.get_quantile_range(), (10.0, 90.0));
}
#[test]
fn invalid_quantile_range_is_rejected() {
for (low, high) in [(75.0, 25.0), (25.0, 25.0), (-1.0, 75.0), (25.0, 101.0)] {
let err = RobustScaler::new()
.with_quantile_range(low, high)
.unwrap_err();
match err {
Error::InvalidParameter { name, .. } => assert_eq!(name, "quantile_range"),
other => panic!("expected InvalidParameter, got {:?}", other),
}
}
}
#[test]
fn constant_feature_maps_to_zeros() {
let x = array![[3.0, 1.0], [3.0, 3.0], [3.0, 5.0]];
let mut scaler = RobustScaler::new();
let z = scaler.fit_transform(&x).unwrap();
assert_eq!(scaler.get_scale().unwrap()[0], 1.0);
assert!(z.column(0).iter().all(|&v| v == 0.0));
assert!(z.iter().all(|v| v.is_finite()));
}
#[test]
fn inverse_transform_round_trips() {
let x = array![[1.0, -5.0], [2.0, 7.5], [3.0, 0.5], [4.0, 100.0]];
let mut scaler = RobustScaler::new();
let z = scaler.fit_transform(&x).unwrap();
let restored = scaler.inverse_transform(&z).unwrap();
for (original, back) in x.iter().zip(restored.iter()) {
assert!((original - back).abs() < 1e-9);
}
}
#[test]
fn flags_control_what_transform_applies() {
let x = array![[1.0], [2.0], [3.0], [4.0]];
let centered = RobustScaler::new()
.with_scaling(false)
.fit_transform(&x)
.unwrap();
assert!((centered[[0, 0]] + 1.5).abs() < 1e-12);
let scaled = RobustScaler::new()
.with_centering(false)
.fit_transform(&x)
.unwrap();
assert!((scaled[[0, 0]] - 1.0 / 1.5).abs() < 1e-12);
}
#[test]
fn single_sample_fit() {
let mut scaler = RobustScaler::new();
scaler.fit(&array![[5.0, -2.0]]).unwrap();
assert_eq!(scaler.get_center().unwrap(), &array![5.0, -2.0]);
assert_eq!(scaler.get_scale().unwrap(), &array![1.0, 1.0]);
}
#[test]
fn transform_before_fit_gives_not_fitted() {
let err = RobustScaler::new().transform(&array![[1.0]]).unwrap_err();
match err {
Error::NotFitted(model) => assert_eq!(model, "RobustScaler"),
other => panic!("expected NotFitted, got {:?}", other),
}
}
}