use core::iter::FromIterator;
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum ColumnSelector {
#[default]
All,
Include(Vec<usize>),
Exclude(Vec<usize>),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct StandardizeParams {
pub with_mean: bool,
pub with_std: bool,
}
impl Default for StandardizeParams {
fn default() -> Self {
Self {
with_mean: true,
with_std: true,
}
}
}
impl StandardizeParams {
#[must_use]
pub const fn with_mean(mut self, with_mean: bool) -> Self {
self.with_mean = with_mean;
self
}
#[must_use]
pub const fn with_std(mut self, with_std: bool) -> Self {
self.with_std = with_std;
self
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum PreprocessingStep {
AddInteractions,
AddPolynomial {
order: usize,
},
ReplaceWithPCA {
number_of_components: usize,
},
ReplaceWithSVD {
number_of_components: usize,
},
Standardize(StandardizeParams),
Scale(ScaleParams),
Impute(ImputeParams),
EncodeCategorical(CategoricalEncoderParams),
PowerTransform(PowerTransformParams),
FilterColumns(ColumnFilterParams),
}
impl core::fmt::Display for PreprocessingStep {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::AddInteractions => write!(f, "Interaction terms added"),
Self::AddPolynomial { order } => {
write!(f, "Polynomial terms added (order = {order})")
}
Self::ReplaceWithPCA {
number_of_components,
} => write!(f, "Replaced with PCA features (n = {number_of_components})"),
Self::ReplaceWithSVD {
number_of_components,
} => write!(f, "Replaced with SVD features (n = {number_of_components})"),
Self::Standardize(params) => write!(
f,
"Standardized features (with_mean = {}, with_std = {})",
params.with_mean, params.with_std
),
Self::Scale(params) => write!(f, "Scaled features using {:?}", params.strategy),
Self::Impute(params) => write!(f, "Imputed features using {:?}", params.strategy),
Self::EncodeCategorical(params) => {
write!(f, "Encoded categorical columns with {:?}", params.encoding)
}
Self::PowerTransform(params) => {
write!(f, "Applied power transform {:?}", params.transform)
}
Self::FilterColumns(params) => {
let mode = if params.retain_selected {
"retain"
} else {
"drop"
};
write!(f, "Column filter ({mode})")
}
}
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ScaleParams {
pub strategy: ScaleStrategy,
pub selector: ColumnSelector,
}
impl Default for ScaleParams {
fn default() -> Self {
Self {
strategy: ScaleStrategy::Standard(StandardizeParams::default()),
selector: ColumnSelector::All,
}
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum ScaleStrategy {
Standard(StandardizeParams),
MinMax(MinMaxParams),
Robust(RobustScaleParams),
}
#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
pub struct MinMaxParams {
pub feature_range: (f64, f64),
}
impl Default for MinMaxParams {
fn default() -> Self {
Self {
feature_range: (0.0, 1.0),
}
}
}
impl MinMaxParams {
#[must_use]
pub const fn with_feature_range(mut self, range: (f64, f64)) -> Self {
self.feature_range = range;
self
}
}
#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
pub struct RobustScaleParams {
pub quantile_range: (f64, f64),
}
impl Default for RobustScaleParams {
fn default() -> Self {
Self {
quantile_range: (25.0, 75.0),
}
}
}
impl RobustScaleParams {
#[must_use]
pub const fn with_quantile_range(mut self, range: (f64, f64)) -> Self {
self.quantile_range = range;
self
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ImputeParams {
pub strategy: ImputeStrategy,
pub selector: ColumnSelector,
}
impl Default for ImputeParams {
fn default() -> Self {
Self {
strategy: ImputeStrategy::Mean,
selector: ColumnSelector::All,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum ImputeStrategy {
Mean,
Median,
MostFrequent,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct CategoricalEncoderParams {
pub selector: ColumnSelector,
pub encoding: CategoricalEncoding,
}
impl Default for CategoricalEncoderParams {
fn default() -> Self {
Self {
selector: ColumnSelector::All,
encoding: CategoricalEncoding::Ordinal,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub enum CategoricalEncoding {
Ordinal,
OneHot {
drop_first: bool,
},
}
impl CategoricalEncoding {
#[must_use]
pub const fn one_hot(drop_first: bool) -> Self {
Self::OneHot { drop_first }
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct PowerTransformParams {
pub selector: ColumnSelector,
pub transform: PowerTransform,
}
impl Default for PowerTransformParams {
fn default() -> Self {
Self {
selector: ColumnSelector::All,
transform: PowerTransform::Log { offset: 0.0 },
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
pub enum PowerTransform {
Log {
offset: f64,
},
BoxCox {
lambda: f64,
},
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ColumnFilterParams {
pub selector: ColumnSelector,
pub retain_selected: bool,
}
impl Default for ColumnFilterParams {
fn default() -> Self {
Self {
selector: ColumnSelector::All,
retain_selected: true,
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct PreprocessingPipeline {
steps: Vec<PreprocessingStep>,
}
impl PreprocessingPipeline {
#[must_use]
pub fn new() -> Self {
Self { steps: Vec::new() }
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.steps.is_empty()
}
#[must_use]
pub fn steps(&self) -> &[PreprocessingStep] {
&self.steps
}
#[must_use]
pub fn add_step(mut self, step: PreprocessingStep) -> Self {
self.steps.push(step);
self
}
pub fn push_step(&mut self, step: PreprocessingStep) {
self.steps.push(step);
}
}
impl From<Vec<PreprocessingStep>> for PreprocessingPipeline {
fn from(steps: Vec<PreprocessingStep>) -> Self {
Self { steps }
}
}
impl From<PreprocessingStep> for PreprocessingPipeline {
fn from(step: PreprocessingStep) -> Self {
Self { steps: vec![step] }
}
}
impl FromIterator<PreprocessingStep> for PreprocessingPipeline {
fn from_iter<T: IntoIterator<Item = PreprocessingStep>>(iter: T) -> Self {
Self {
steps: iter.into_iter().collect(),
}
}
}
impl core::fmt::Display for PreprocessingPipeline {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
if self.steps.is_empty() {
return write!(f, "No preprocessing");
}
for (idx, step) in self.steps.iter().enumerate() {
if idx > 0 {
write!(f, " -> ")?;
}
write!(f, "{step}")?;
}
Ok(())
}
}