use super::{fitted, for_each_row, validate_matrix, validate_transform_matrix};
use crate::error::Error;
use crate::parallel_gates::scan_f64_parallel_min_elems;
use crate::utils::standardize::{WelfordState, scale_from_variance, welford_merge, welford_step};
use crate::{Deserialize, Serialize};
use ndarray::{Array1, Array2, ArrayBase, ArrayView1, ArrayView2, Axis, Data, Ix2};
use rayon::iter::{IntoParallelIterator, ParallelIterator};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StandardScaler {
with_mean: bool,
with_std: bool,
mean: Option<Array1<f64>>,
var: Option<Array1<f64>>,
scale: Option<Array1<f64>>,
n_samples_seen: usize,
}
impl Default for StandardScaler {
fn default() -> Self {
Self::new()
}
}
impl StandardScaler {
pub fn new() -> Self {
Self {
with_mean: true,
with_std: true,
mean: None,
var: None,
scale: None,
n_samples_seen: 0,
}
}
pub fn with_mean(mut self, with_mean: bool) -> Self {
self.with_mean = with_mean;
self
}
pub fn with_std(mut self, with_std: bool) -> Self {
self.with_std = with_std;
self
}
get_field!(get_with_mean, with_mean, bool);
get_field!(get_with_std, with_std, bool);
get_field!(get_n_samples_seen, n_samples_seen, usize);
get_field_as_ref!(get_mean, mean, Option<&Array1<f64>>);
get_field_as_ref!(get_var, var, Option<&Array1<f64>>);
get_field_as_ref!(get_scale, scale, Option<&Array1<f64>>);
#[inline]
pub fn get_n_features(&self) -> Option<usize> {
self.mean.as_ref().map(|mean| mean.len())
}
pub fn fit<S>(&mut self, x: &ArrayBase<S, Ix2>) -> Result<&mut Self, Error>
where
S: Data<Elem = f64>,
{
validate_matrix(x)?;
let states = column_states(&x.view());
self.store(&states);
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 states = column_states(&x.view());
if let (Some(mean), Some(var)) = (&self.mean, &self.var) {
let seen = self.n_samples_seen as f64;
for (j, state) in states.iter_mut().enumerate() {
*state = welford_merge((seen, mean[j], var[j] * seen), *state);
}
}
self.store(&states);
Ok(self)
}
pub fn transform<S>(&self, x: &ArrayBase<S, Ix2>) -> Result<Array2<f64>, Error>
where
S: Data<Elem = f64>,
{
let (mean, scale) = self.fitted_stats()?;
validate_transform_matrix(x, mean.len())?;
let mut result = x.to_owned();
let (with_mean, with_std) = (self.with_mean, self.with_std);
for_each_row(&mut result, |mut row| match (with_mean, with_std) {
(true, true) => {
for ((value, &m), &s) in row.iter_mut().zip(mean).zip(scale) {
*value = (*value - m) / s;
}
}
(true, false) => {
for (value, &m) in row.iter_mut().zip(mean) {
*value -= m;
}
}
(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 (mean, scale) = self.fitted_stats()?;
validate_transform_matrix(x, mean.len())?;
let mut result = x.to_owned();
let (with_mean, with_std) = (self.with_mean, self.with_std);
for_each_row(&mut result, |mut row| match (with_mean, with_std) {
(true, true) => {
for ((value, &m), &s) in row.iter_mut().zip(mean).zip(scale) {
*value = *value * s + m;
}
}
(true, false) => {
for (value, &m) in row.iter_mut().zip(mean) {
*value += m;
}
}
(false, true) => {
for (value, &s) in row.iter_mut().zip(scale) {
*value *= s;
}
}
(false, false) => {}
});
Ok(result)
}
model_save_and_load_methods!(StandardScaler);
fn fitted_stats(&self) -> Result<(&Array1<f64>, &Array1<f64>), Error> {
Ok((
fitted(&self.mean, "StandardScaler")?,
fitted(&self.scale, "StandardScaler")?,
))
}
fn store(&mut self, states: &[WelfordState]) {
let n = states.first().map_or(0.0, |&(count, _, _)| count);
self.mean = Some(states.iter().map(|&(_, mean, _)| mean).collect());
self.var = Some(states.iter().map(|&(_, _, m2)| m2 / n).collect());
self.scale = Some(
states
.iter()
.map(|&(_, mean, m2)| scale_from_variance(m2 / n, mean, n))
.collect(),
);
self.n_samples_seen = n as usize;
}
}
fn column_states(x: &ArrayView2<f64>) -> Vec<WelfordState> {
let fold_lane = |lane: ArrayView1<f64>| {
lane.iter()
.fold((0.0, 0.0, 0.0), |acc, &value| welford_step(acc, value))
};
let lanes: Vec<ArrayView1<f64>> = x.lanes(Axis(0)).into_iter().collect();
if x.len() >= scan_f64_parallel_min_elems() {
lanes.into_par_iter().map(fold_lane).collect()
} else {
lanes.into_iter().map(fold_lane).collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::utils::standardize::{StandardizationAxis, standardize};
use ndarray::array;
#[test]
fn fit_transform_matches_standardize_column() {
let x = array![[1.0, 2000.0], [2.0, 3000.0], [3.0, 4000.0], [4.0, 5000.0]];
let scaled = StandardScaler::new().fit_transform(&x).unwrap();
let stateless = standardize(&x, StandardizationAxis::Column).unwrap();
assert_eq!(scaled, stateless);
}
#[test]
fn transform_reuses_training_statistics() {
let x_train = array![[1.0], [2.0], [3.0]];
let mut scaler = StandardScaler::new();
scaler.fit(&x_train).unwrap();
let z = scaler.transform(&array![[3.0]]).unwrap();
let expected = (3.0 - 2.0) / (2.0f64 / 3.0).sqrt();
assert!((z[[0, 0]] - expected).abs() < 1e-12);
assert!(z[[0, 0]] > 1.0);
}
#[test]
fn partial_fit_matches_single_fit() {
let batch_a = array![[1.0, 10.0], [2.0, 20.0]];
let batch_b = array![[3.0, 30.0], [4.0, 40.0], [5.0, 50.0]];
let full = array![
[1.0, 10.0],
[2.0, 20.0],
[3.0, 30.0],
[4.0, 40.0],
[5.0, 50.0]
];
let mut incremental = StandardScaler::new();
incremental.partial_fit(&batch_a).unwrap();
incremental.partial_fit(&batch_b).unwrap();
let mut single = StandardScaler::new();
single.fit(&full).unwrap();
assert_eq!(incremental.get_n_samples_seen(), 5);
for j in 0..2 {
assert!(
(incremental.get_mean().unwrap()[j] - single.get_mean().unwrap()[j]).abs() < 1e-9
);
assert!(
(incremental.get_var().unwrap()[j] - single.get_var().unwrap()[j]).abs() < 1e-9
);
}
}
#[test]
fn inverse_transform_round_trips() {
let x = array![[1.0, -5.0], [2.0, 7.5], [3.0, 0.5]];
let mut scaler = StandardScaler::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 constant_feature_maps_to_zeros() {
let x = array![[3.0, 1.0], [3.0, 3.0], [3.0, 5.0]];
let mut scaler = StandardScaler::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.column(1).iter().all(|v| v.is_finite()));
}
#[test]
fn flags_control_what_transform_applies() {
let x = array![[1.0], [2.0], [3.0]];
let centered_only = StandardScaler::new()
.with_std(false)
.fit_transform(&x)
.unwrap();
assert_eq!(centered_only, array![[-1.0], [0.0], [1.0]]);
let scaled_only = StandardScaler::new()
.with_mean(false)
.fit_transform(&x)
.unwrap();
let std = (2.0f64 / 3.0).sqrt();
assert!((scaled_only[[0, 0]] - 1.0 / std).abs() < 1e-12);
}
#[test]
fn transform_before_fit_gives_not_fitted() {
let err = StandardScaler::new()
.transform(&array![[1.0, 2.0]])
.unwrap_err();
match err {
Error::NotFitted(model) => assert_eq!(model, "StandardScaler"),
other => panic!("expected NotFitted, got {:?}", other),
}
}
#[test]
fn transform_feature_mismatch_gives_dimension_mismatch() {
let mut scaler = StandardScaler::new();
scaler.fit(&array![[1.0, 2.0], [3.0, 4.0]]).unwrap();
let err = scaler.transform(&array![[1.0, 2.0, 3.0]]).unwrap_err();
match err {
Error::DimensionMismatch { expected, found } => {
assert_eq!(expected, 2);
assert_eq!(found, 3);
}
other => panic!("expected DimensionMismatch, got {:?}", other),
}
}
#[test]
fn non_finite_input_is_rejected() {
let err = StandardScaler::new()
.fit(&array![[1.0, f64::NAN], [3.0, 4.0]])
.unwrap_err();
match err {
Error::NonFinite(_) => {}
other => panic!("expected NonFinite, got {:?}", other),
}
}
}