Skip to main content

millwright/
balance.rs

1//! Train-time balancers — the resampling stage of a pipeline.
2//!
3//! These adapt [`imbalance-rs`](https://docs.rs/imbalance-rs) samplers behind
4//! the framework's [`Balancer`] trait. A balancer runs **only during `fit`**
5//! (via [`Pipeline::balance`](crate::pipeline::Pipeline::balance)) — never at
6//! predict time — because resampling changes the row set.
7//!
8//! Conversion to `ndarray` happens here, at the edge, so the imbalance-rs
9//! array world never reaches user code.
10
11use imbalance_rs::{RandomOverSampler as ImbRos, Sampler, Smote as ImbSmote};
12use ndarray::{Array1, Array2};
13
14use crate::error::{Error, Result};
15use crate::frame::Frame;
16use crate::traits::Balancer;
17
18fn frame_to_array2(frame: &Frame) -> Result<Array2<f64>> {
19    let (n, p) = frame.shape();
20    Array2::from_shape_vec((n, p), frame.buf().to_vec())
21        .map_err(|e| Error::Backend(format!("ndarray conversion failed: {e}")))
22}
23
24fn array2_to_frame(arr: &Array2<f64>, columns: &[String]) -> Result<Frame> {
25    let (n, p) = arr.dim();
26    // `.iter()` yields elements in row-major logical order, which is exactly
27    // the layout `Frame` stores.
28    let buf: Vec<f64> = arr.iter().copied().collect();
29    Frame::new(buf, n, p, columns.to_vec())
30}
31
32/// SMOTE over-sampling (Synthetic Minority Over-sampling Technique).
33///
34/// Synthesises new minority-class rows by interpolating between a sample and its
35/// nearest same-class neighbours. Backed by `imbalance_rs::Smote`.
36#[derive(Clone, Debug)]
37pub struct Smote {
38    k_neighbors: usize,
39    seed: u64,
40}
41
42impl Smote {
43    /// SMOTE with the default 5 neighbours.
44    pub fn new() -> Self {
45        Smote {
46            k_neighbors: 5,
47            seed: 0,
48        }
49    }
50
51    /// Number of nearest neighbours used to interpolate.
52    pub fn k_neighbors(mut self, k: usize) -> Self {
53        self.k_neighbors = k;
54        self
55    }
56
57    /// Seed the RNG for reproducible synthesis.
58    pub fn random_state(mut self, seed: u64) -> Self {
59        self.seed = seed;
60        self
61    }
62}
63
64impl Default for Smote {
65    fn default() -> Self {
66        Smote::new()
67    }
68}
69
70impl Balancer for Smote {
71    fn name(&self) -> &'static str {
72        "Smote"
73    }
74
75    fn fit_resample(&self, features: &Frame, target: &[f64]) -> Result<(Frame, Vec<f64>)> {
76        let x = frame_to_array2(features)?;
77        let y: Array1<i64> =
78            Array1::from(target.iter().map(|v| v.round() as i64).collect::<Vec<_>>());
79        let sampler = ImbSmote::new()
80            .k_neighbors(self.k_neighbors)
81            .random_state(self.seed);
82        let (xr, yr) = sampler
83            .fit_resample(&x, &y)
84            .map_err(|e| Error::Backend(format!("SMOTE failed: {e}")))?;
85        let frame = array2_to_frame(&xr, features.columns())?;
86        let target = yr.iter().map(|&l| l as f64).collect();
87        Ok((frame, target))
88    }
89}
90
91/// Random over-sampling: duplicate minority-class rows until balanced.
92///
93/// Backed by `imbalance_rs::RandomOverSampler`.
94#[derive(Clone, Debug, Default)]
95pub struct RandomOverSampler;
96
97impl RandomOverSampler {
98    pub fn new() -> Self {
99        RandomOverSampler
100    }
101}
102
103impl Balancer for RandomOverSampler {
104    fn name(&self) -> &'static str {
105        "RandomOverSampler"
106    }
107
108    fn fit_resample(&self, features: &Frame, target: &[f64]) -> Result<(Frame, Vec<f64>)> {
109        let x = frame_to_array2(features)?;
110        let y: Array1<i64> =
111            Array1::from(target.iter().map(|v| v.round() as i64).collect::<Vec<_>>());
112        let (xr, yr) = ImbRos::new()
113            .fit_resample(&x, &y)
114            .map_err(|e| Error::Backend(format!("RandomOverSampler failed: {e}")))?;
115        let frame = array2_to_frame(&xr, features.columns())?;
116        let target = yr.iter().map(|&l| l as f64).collect();
117        Ok((frame, target))
118    }
119}
120
121#[cfg(test)]
122mod tests {
123    use super::*;
124
125    #[test]
126    fn smote_balances_the_minority_class() {
127        // 6 majority (class 0), 2 minority (class 1).
128        let x = Frame::from_rows(
129            vec![
130                vec![0.0, 0.0],
131                vec![0.1, 0.2],
132                vec![0.2, 0.1],
133                vec![0.3, 0.0],
134                vec![0.0, 0.3],
135                vec![0.2, 0.2],
136                vec![9.0, 9.0],
137                vec![9.1, 9.2],
138            ],
139            vec!["a".into(), "b".into()],
140        )
141        .unwrap();
142        let y = vec![0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0, 1.0];
143
144        let (xr, yr) = Smote::new()
145            .k_neighbors(1)
146            .random_state(7)
147            .fit_resample(&x, &y)
148            .unwrap();
149        let zeros = yr.iter().filter(|&&v| v == 0.0).count();
150        let ones = yr.iter().filter(|&&v| v == 1.0).count();
151        assert_eq!(zeros, 6);
152        assert_eq!(ones, 6, "minority class should be oversampled to parity");
153        assert_eq!(xr.nrows(), yr.len());
154        assert_eq!(xr.ncols(), 2);
155    }
156}