use super::{
column_min_max, 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, Default, Serialize, Deserialize)]
pub struct MaxAbsScaler {
max_abs: Option<Array1<f64>>,
scale: Option<Array1<f64>>,
n_samples_seen: usize,
}
impl MaxAbsScaler {
pub fn new() -> Self {
Self::default()
}
get_field!(get_n_samples_seen, n_samples_seen, usize);
get_field_as_ref!(get_max_abs, max_abs, Option<&Array1<f64>>);
get_field_as_ref!(get_scale, scale, Option<&Array1<f64>>);
#[inline]
pub fn get_n_features(&self) -> Option<usize> {
self.max_abs.as_ref().map(|max_abs| max_abs.len())
}
pub fn fit<S>(&mut self, x: &ArrayBase<S, Ix2>) -> Result<&mut Self, Error>
where
S: Data<Elem = f64>,
{
validate_matrix(x)?;
self.max_abs = Some(column_max_abs(x));
self.n_samples_seen = x.nrows();
self.recompute_scale();
Ok(self)
}
pub fn partial_fit<S>(&mut self, x: &ArrayBase<S, Ix2>) -> Result<&mut Self, Error>
where
S: Data<Elem = f64>,
{
validate_matrix(x)?;
if let Some(n_features) = self.get_n_features()
&& n_features != x.ncols()
{
return Err(Error::dimension_mismatch(n_features, x.ncols()));
}
let mut max_abs = column_max_abs(x);
if let Some(seen) = &self.max_abs {
max_abs.zip_mut_with(seen, |value, &seen| *value = value.max(seen));
}
self.max_abs = Some(max_abs);
self.n_samples_seen += x.nrows();
self.recompute_scale();
Ok(self)
}
pub fn transform<S>(&self, x: &ArrayBase<S, Ix2>) -> Result<Array2<f64>, Error>
where
S: Data<Elem = f64>,
{
let scale = fitted(&self.scale, "MaxAbsScaler")?;
validate_transform_matrix(x, scale.len())?;
let mut result = x.to_owned();
for_each_row(&mut result, |mut row| {
for (value, &s) in row.iter_mut().zip(scale) {
*value /= s;
}
});
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 scale = fitted(&self.scale, "MaxAbsScaler")?;
validate_transform_matrix(x, scale.len())?;
let mut result = x.to_owned();
for_each_row(&mut result, |mut row| {
for (value, &s) in row.iter_mut().zip(scale) {
*value *= s;
}
});
Ok(result)
}
model_save_and_load_methods!(MaxAbsScaler);
fn recompute_scale(&mut self) {
self.scale = self
.max_abs
.as_ref()
.map(|max_abs| max_abs.mapv(handle_zero_scale));
}
}
fn column_max_abs<S>(x: &ArrayBase<S, Ix2>) -> Array1<f64>
where
S: Data<Elem = f64>,
{
column_min_max(&x.view())
.into_iter()
.map(|(lo, hi)| lo.abs().max(hi.abs()))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
#[test]
fn divides_by_column_magnitude() {
let x = array![[1.0, -4.0], [0.0, 2.0], [-2.0, 0.0]];
let mut scaler = MaxAbsScaler::new();
let z = scaler.fit_transform(&x).unwrap();
assert_eq!(scaler.get_max_abs().unwrap(), &array![2.0, 4.0]);
assert_eq!(z, array![[0.5, -1.0], [0.0, 0.5], [-1.0, 0.0]]);
assert_eq!(scaler.get_n_samples_seen(), 3);
}
#[test]
fn training_values_stay_within_unit_magnitude() {
let x = array![[10.0, -3.0], [-7.0, 1.0], [4.0, 2.0]];
let z = MaxAbsScaler::new().fit_transform(&x).unwrap();
assert!(z.iter().all(|v| v.abs() <= 1.0));
}
#[test]
fn all_zero_feature_stays_zero() {
let x = array![[0.0, 1.0], [0.0, 2.0]];
let mut scaler = MaxAbsScaler::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));
}
#[test]
fn partial_fit_keeps_the_larger_magnitude() {
let batch_a = array![[1.0, -9.0]];
let batch_b = array![[-5.0, 2.0]];
let full = array![[1.0, -9.0], [-5.0, 2.0]];
let mut incremental = MaxAbsScaler::new();
incremental.partial_fit(&batch_a).unwrap();
incremental.partial_fit(&batch_b).unwrap();
let mut single = MaxAbsScaler::new();
single.fit(&full).unwrap();
assert_eq!(incremental.get_max_abs(), single.get_max_abs());
assert_eq!(incremental.get_n_samples_seen(), 2);
}
#[test]
fn inverse_transform_round_trips() {
let x = array![[1.0, -5.0], [2.0, 7.5], [0.0, 0.5]];
let mut scaler = MaxAbsScaler::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 transform_before_fit_gives_not_fitted() {
let err = MaxAbsScaler::new().transform(&array![[1.0]]).unwrap_err();
match err {
Error::NotFitted(model) => assert_eq!(model, "MaxAbsScaler"),
other => panic!("expected NotFitted, got {:?}", other),
}
}
}