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, Serialize, Deserialize)]
pub struct MinMaxScaler {
feature_range: (f64, f64),
clip: bool,
data_min: Option<Array1<f64>>,
data_max: Option<Array1<f64>>,
scale: Option<Array1<f64>>,
min: Option<Array1<f64>>,
n_samples_seen: usize,
}
impl Default for MinMaxScaler {
fn default() -> Self {
Self::new()
}
}
impl MinMaxScaler {
pub fn new() -> Self {
Self {
feature_range: (0.0, 1.0),
clip: false,
data_min: None,
data_max: None,
scale: None,
min: None,
n_samples_seen: 0,
}
}
pub fn with_feature_range(mut self, min: f64, max: f64) -> Result<Self, Error> {
if !min.is_finite() || !max.is_finite() || min >= max {
return Err(Error::invalid_parameter(
"feature_range",
format!("must be finite with min < max, got ({}, {})", min, max),
));
}
self.feature_range = (min, max);
self.recompute_affine();
Ok(self)
}
pub fn with_clip(mut self, clip: bool) -> Self {
self.clip = clip;
self
}
get_field!(get_feature_range, feature_range, (f64, f64));
get_field!(get_clip, clip, bool);
get_field!(get_n_samples_seen, n_samples_seen, usize);
get_field_as_ref!(get_data_min, data_min, Option<&Array1<f64>>);
get_field_as_ref!(get_data_max, data_max, Option<&Array1<f64>>);
get_field_as_ref!(get_scale, scale, Option<&Array1<f64>>);
get_field_as_ref!(get_min, min, Option<&Array1<f64>>);
pub fn get_data_range(&self) -> Option<Array1<f64>> {
match (&self.data_min, &self.data_max) {
(Some(min), Some(max)) => Some(max - min),
_ => None,
}
}
#[inline]
pub fn get_n_features(&self) -> Option<usize> {
self.data_min.as_ref().map(|min| min.len())
}
pub fn fit<S>(&mut self, x: &ArrayBase<S, Ix2>) -> Result<&mut Self, Error>
where
S: Data<Elem = f64>,
{
validate_matrix(x)?;
let extrema = column_min_max(&x.view());
self.data_min = Some(extrema.iter().map(|&(lo, _)| lo).collect());
self.data_max = Some(extrema.iter().map(|&(_, hi)| hi).collect());
self.n_samples_seen = x.nrows();
self.recompute_affine();
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 extrema = column_min_max(&x.view());
let mut data_min: Array1<f64> = extrema.iter().map(|&(lo, _)| lo).collect();
let mut data_max: Array1<f64> = extrema.iter().map(|&(_, hi)| hi).collect();
if let (Some(seen_min), Some(seen_max)) = (&self.data_min, &self.data_max) {
data_min.zip_mut_with(seen_min, |value, &seen| *value = value.min(seen));
data_max.zip_mut_with(seen_max, |value, &seen| *value = value.max(seen));
}
self.data_min = Some(data_min);
self.data_max = Some(data_max);
self.n_samples_seen += x.nrows();
self.recompute_affine();
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, "MinMaxScaler")?;
let offset = fitted(&self.min, "MinMaxScaler")?;
validate_transform_matrix(x, scale.len())?;
let mut result = x.to_owned();
let (low, high) = self.feature_range;
let clip = self.clip;
for_each_row(&mut result, |mut row| {
for ((value, &s), &o) in row.iter_mut().zip(scale).zip(offset) {
let scaled = *value * s + o;
*value = if clip {
scaled.clamp(low, high)
} else {
scaled
};
}
});
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, "MinMaxScaler")?;
let offset = fitted(&self.min, "MinMaxScaler")?;
validate_transform_matrix(x, scale.len())?;
let mut result = x.to_owned();
for_each_row(&mut result, |mut row| {
for ((value, &s), &o) in row.iter_mut().zip(scale).zip(offset) {
*value = (*value - o) / s;
}
});
Ok(result)
}
model_save_and_load_methods!(MinMaxScaler);
fn recompute_affine(&mut self) {
let (Some(data_min), Some(data_max)) = (&self.data_min, &self.data_max) else {
return;
};
let (low, high) = self.feature_range;
let span = high - low;
let scale: Array1<f64> = data_min
.iter()
.zip(data_max)
.map(|(&lo, &hi)| span / handle_zero_scale(hi - lo))
.collect();
let min: Array1<f64> = data_min
.iter()
.zip(&scale)
.map(|(&lo, &s)| low - lo * s)
.collect();
self.scale = Some(scale);
self.min = Some(min);
}
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
#[test]
fn default_range_maps_training_extrema_to_unit_interval() {
let x = array![[1.0, 10.0], [2.0, 20.0], [3.0, 30.0]];
let mut scaler = MinMaxScaler::new();
let z = scaler.fit_transform(&x).unwrap();
assert_eq!(z, array![[0.0, 0.0], [0.5, 0.5], [1.0, 1.0]]);
assert_eq!(scaler.get_data_min().unwrap(), &array![1.0, 10.0]);
assert_eq!(scaler.get_data_max().unwrap(), &array![3.0, 30.0]);
assert_eq!(scaler.get_data_range().unwrap(), array![2.0, 20.0]);
}
#[test]
fn custom_feature_range() {
let x = array![[0.0], [10.0]];
let mut scaler = MinMaxScaler::new().with_feature_range(-1.0, 1.0).unwrap();
let z = scaler.fit_transform(&x).unwrap();
assert_eq!(z, array![[-1.0], [1.0]]);
assert_eq!(scaler.get_feature_range(), (-1.0, 1.0));
}
#[test]
fn invalid_feature_range_is_rejected() {
for (low, high) in [(1.0, 1.0), (2.0, 1.0), (f64::NAN, 1.0)] {
let err = MinMaxScaler::new()
.with_feature_range(low, high)
.unwrap_err();
match err {
Error::InvalidParameter { name, .. } => assert_eq!(name, "feature_range"),
other => panic!("expected InvalidParameter, got {:?}", other),
}
}
}
#[test]
fn clip_bounds_out_of_range_values() {
let x_train = array![[1.0], [3.0]];
let mut open = MinMaxScaler::new();
open.fit(&x_train).unwrap();
assert_eq!(open.transform(&array![[5.0]]).unwrap(), array![[2.0]]);
let mut clipped = MinMaxScaler::new().with_clip(true);
clipped.fit(&x_train).unwrap();
assert_eq!(clipped.transform(&array![[5.0]]).unwrap(), array![[1.0]]);
assert_eq!(clipped.transform(&array![[-3.0]]).unwrap(), array![[0.0]]);
}
#[test]
fn constant_feature_maps_to_range_start() {
let x = array![[3.0, 1.0], [3.0, 3.0]];
let mut scaler = MinMaxScaler::new().with_feature_range(-2.0, 2.0).unwrap();
let z = scaler.fit_transform(&x).unwrap();
assert!(z.iter().all(|v| v.is_finite()));
assert_eq!(z.column(0).to_vec(), vec![-2.0, -2.0]);
}
#[test]
fn partial_fit_widens_the_interval() {
let batch_a = array![[2.0, 5.0], [3.0, 4.0]];
let batch_b = array![[1.0, 9.0]];
let full = array![[2.0, 5.0], [3.0, 4.0], [1.0, 9.0]];
let mut incremental = MinMaxScaler::new();
incremental.partial_fit(&batch_a).unwrap();
incremental.partial_fit(&batch_b).unwrap();
let mut single = MinMaxScaler::new();
single.fit(&full).unwrap();
assert_eq!(incremental.get_data_min(), single.get_data_min());
assert_eq!(incremental.get_data_max(), single.get_data_max());
assert_eq!(incremental.get_n_samples_seen(), 3);
}
#[test]
fn inverse_transform_round_trips() {
let x = array![[1.0, -5.0], [2.0, 7.5], [3.0, 0.5]];
let mut scaler = MinMaxScaler::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 = MinMaxScaler::new().transform(&array![[1.0]]).unwrap_err();
match err {
Error::NotFitted(model) => assert_eq!(model, "MinMaxScaler"),
other => panic!("expected NotFitted, got {:?}", other),
}
}
}